diff --git a/packages/ai/example/tutorial.ts b/packages/ai/example/tutorial.ts index b109ef6230..681a4b42f9 100644 --- a/packages/ai/example/tutorial.ts +++ b/packages/ai/example/tutorial.ts @@ -194,7 +194,7 @@ const FakeProtocol = Protocol.make({ event: Schema.String, initial: () => undefined, step: (_, frame) => Effect.succeed([undefined, [{ type: "text-delta", id: "text-0", text: frame }]] as const), - onHalt: () => [{ type: "finish", reason: "stop" }], + onHalt: () => Effect.succeed([{ type: "finish", reason: "stop" }]), }, }) diff --git a/packages/ai/src/protocols/bedrock-converse.ts b/packages/ai/src/protocols/bedrock-converse.ts index b94f23d62f..643e569527 100644 --- a/packages/ai/src/protocols/bedrock-converse.ts +++ b/packages/ai/src/protocols/bedrock-converse.ts @@ -614,7 +614,7 @@ const step = (state: ParserState, event: BedrockEvent) => const framing = BedrockEventStream.framing(ADAPTER) -const onHalt = (state: ParserState): ReadonlyArray => +const finishEvents = (state: ParserState): ReadonlyArray => state.pendingFinish ? (() => { const events: LLMEvent[] = [] @@ -650,7 +650,7 @@ export const protocol = Protocol.make({ reasoningSignatures: {}, }), step, - onHalt, + onHalt: (state) => Effect.succeed(finishEvents(state)), }, }) diff --git a/packages/ai/src/protocols/gemini.ts b/packages/ai/src/protocols/gemini.ts index 82a69b059e..f9f03aae75 100644 --- a/packages/ai/src/protocols/gemini.ts +++ b/packages/ai/src/protocols/gemini.ts @@ -376,7 +376,7 @@ const mapFinishReason = (finishReason: string | undefined, hasToolCalls: boolean return "unknown" } -const finish = (state: ParserState): ReadonlyArray => +const finishEvents = (state: ParserState): ReadonlyArray => state.finishReason || state.usage ? (() => { const events: LLMEvent[] = [] @@ -493,7 +493,7 @@ export const protocol = Protocol.make({ event: Protocol.jsonEvent(GeminiEvent), initial: () => ({ hasToolCalls: false, nextToolCallId: 0, lifecycle: Lifecycle.initial() }), step, - onHalt: finish, + onHalt: (state) => Effect.succeed(finishEvents(state)), }, }) diff --git a/packages/ai/src/protocols/openai-chat.ts b/packages/ai/src/protocols/openai-chat.ts index 3b35bb88d8..16900c7b69 100644 --- a/packages/ai/src/protocols/openai-chat.ts +++ b/packages/ai/src/protocols/openai-chat.ts @@ -599,10 +599,14 @@ const step = (state: ParserState, event: OpenAIChatEvent) => ] as const }) -const finishEvents = (state: ParserState): ReadonlyArray => { +const finishEvents = Effect.fn("OpenAIChat.finishEvents")(function* (state: ParserState) { + if (Object.keys(state.pendingTools).length > 0) + return yield* ProviderShared.eventError(ADAPTER, "OpenAI Chat tool call delta is missing id or name") + const finished = Object.keys(state.tools).length > 0 ? yield* ToolStream.finishAll(ADAPTER, state.tools) : undefined + const toolCallEvents = finished?.events ?? state.toolCallEvents const events: LLMEvent[] = [] - const hasToolCalls = state.toolCallEvents.length > 0 - const reason = state.finishReason === "stop" && hasToolCalls ? "tool-calls" : state.finishReason + const hasToolCalls = toolCallEvents.some(LLMEvent.is.toolCall) + const reason = state.finishReason === "stop" && hasToolCalls ? "tool-calls" : (state.finishReason ?? "unknown") const metadata = reasoningMetadata( state.reasoningField, state.reasoningDetailsObserved ? state.reasoningDetails : undefined, @@ -612,11 +616,11 @@ const finishEvents = (state: ParserState): ReadonlyArray => { ? Lifecycle.reasoningStart(state.lifecycle, events, "reasoning-0", reasoningMetadata(state.reasoningField)) : state.lifecycle const ended = Lifecycle.reasoningEnd(started, events, "reasoning-0", metadata) - const lifecycle = state.toolCallEvents.length ? Lifecycle.stepStart(ended, events) : ended - events.push(...state.toolCallEvents) - if (reason) Lifecycle.finish(lifecycle, events, { reason, usage: state.usage }) + const lifecycle = toolCallEvents.length ? Lifecycle.stepStart(ended, events) : ended + events.push(...toolCallEvents) + Lifecycle.finish(lifecycle, events, { reason, usage: state.usage }) return events -} +}) // ============================================================================= // Protocol And OpenAI Route diff --git a/packages/ai/src/route/client.ts b/packages/ai/src/route/client.ts index 067292329b..6d54be5218 100644 --- a/packages/ai/src/route/client.ts +++ b/packages/ai/src/route/client.ts @@ -252,6 +252,33 @@ const requireTerminalEvent = (route: string) => (events: Stream.Stream( + protocol: Protocol, + request: LLMRequest, + events: Stream.Stream, +) => + Stream.suspend(() => { + let state = protocol.stream.initial(request) + const parsed = events.pipe( + Stream.mapEffect((event) => + protocol.stream.step(state, event).pipe( + Effect.map(([next, output]) => { + state = next + return output + }), + ), + ), + Stream.flatMap(Stream.fromIterable), + ) + const onHalt = protocol.stream.onHalt + if (!onHalt) return parsed + return parsed.pipe( + Stream.concat( + Stream.fromEffect(Effect.suspend(() => onHalt(state))).pipe(Stream.flatMap(Stream.fromIterable)), + ), + ) + }) + function makeFromTransport( input: MakeTransportInput, ): Route { @@ -314,12 +341,7 @@ function makeFromTransport( Stream.mapEffect(decodeEvent(route)), protocol.stream.terminal ? Stream.takeUntil(protocol.stream.terminal) : (stream) => stream, ) - return events.pipe( - Stream.mapAccumEffect( - () => protocol.stream.initial(request), - protocol.stream.step, - protocol.stream.onHalt ? { onHalt: protocol.stream.onHalt } : undefined, - ), + return parseEvents(protocol, request, events).pipe( Stream.catchCause((cause) => Stream.fail(streamError(route, `Failed to read ${route} stream`, cause))), requireTerminalEvent(route), ) diff --git a/packages/ai/src/route/protocol.ts b/packages/ai/src/route/protocol.ts index acb1e78c67..76588d724f 100644 --- a/packages/ai/src/route/protocol.ts +++ b/packages/ai/src/route/protocol.ts @@ -58,8 +58,8 @@ export interface ProtocolStream { readonly step: (state: State, event: Event) => Effect.Effect], LLMError> /** Optional request-completion signal for transports that do not end naturally. */ readonly terminal?: (event: Event) => boolean - /** Optional flush emitted when the framed stream ends. */ - readonly onHalt?: (state: State) => ReadonlyArray + /** Optional effectful flush emitted when the framed stream ends. */ + readonly onHalt?: (state: State) => Effect.Effect, LLMError> } /** diff --git a/packages/ai/test/provider/openai-chat.test.ts b/packages/ai/test/provider/openai-chat.test.ts index 591e5f0be0..1849d028ab 100644 --- a/packages/ai/test/provider/openai-chat.test.ts +++ b/packages/ai/test/provider/openai-chat.test.ts @@ -1036,7 +1036,38 @@ describe("OpenAI Chat route", () => { }), ) - it.effect("fails a streamed tool call when the provider ends without a finish reason", () => + it.effect("finishes text with an unknown reason when the provider cleanly ends without one", () => + Effect.gen(function* () { + const response = yield* LLMClient.generate(request).pipe( + Effect.provide(fixedResponse(sseEvents(deltaChunk({ role: "assistant", content: "Hello" })))), + ) + + expect(response.text).toBe("Hello") + expect(response.finishReason).toBe("unknown") + }), + ) + + it.effect("rejects missing tool identity when the provider cleanly ends", () => + Effect.gen(function* () { + for (const tool of [ + { index: 0, id: "call_1", function: { arguments: "{}" } }, + { index: 0, function: { name: "lookup", arguments: "{}" } }, + ]) { + const error = yield* LLMClient.generate( + LLM.updateRequest(request, { + tools: [{ name: "lookup", description: "Lookup data", inputSchema: { type: "object" } }], + }), + ).pipe( + Effect.provide(fixedResponse(sseEvents(deltaChunk({ tool_calls: [tool] })))), + Effect.flip, + ) + + expect(error.message).toContain("OpenAI Chat tool call delta is missing id or name") + } + }), + ) + + it.effect("finalizes a streamed tool call when the provider ends without a finish reason", () => Effect.gen(function* () { const body = sseEvents( deltaChunk({ @@ -1049,23 +1080,56 @@ describe("OpenAI Chat route", () => { tools: [{ name: "lookup", description: "Lookup data", inputSchema: { type: "object" } }], }) const events: LLMEvent[] = [] - const streamError = yield* LLMClient.stream(input).pipe( + yield* LLMClient.stream(input).pipe( Stream.runForEach((event) => Effect.sync(() => events.push(event))), - Effect.flip, Effect.provide(fixedResponse(body)), ) - const error = yield* LLMClient.generate(input).pipe(Effect.provide(fixedResponse(body)), Effect.flip) + const response = yield* LLMClient.generate(input).pipe(Effect.provide(fixedResponse(body))) expect(events).toEqual([ { type: "step-start", index: 0 }, { type: "tool-input-start", id: "call_1", name: "lookup", providerMetadata: undefined }, { type: "tool-input-delta", id: "call_1", name: "lookup", text: '{"query"' }, { type: "tool-input-delta", id: "call_1", name: "lookup", text: ':"weather"}' }, + { type: "tool-input-end", id: "call_1", name: "lookup", providerMetadata: undefined }, + { + type: "tool-call", + id: "call_1", + name: "lookup", + input: { query: "weather" }, + providerExecuted: undefined, + providerMetadata: undefined, + }, + { type: "step-finish", index: 0, reason: "unknown", usage: undefined, providerMetadata: undefined }, + { type: "finish", reason: "unknown", usage: undefined }, ]) - expect(events.filter(LLMEvent.is.toolCall)).toEqual([]) - expect(streamError.reason).toMatchObject({ _tag: "InvalidProviderOutput" }) - expect(streamError.message).toContain("Provider stream ended without a terminal finish event") - expect(error.message).toContain("Provider stream ended without a terminal finish event") + expect(response.toolCalls).toMatchObject([{ id: "call_1", name: "lookup", input: { query: "weather" } }]) + expect(response.finishReason).toBe("unknown") + }), + ) + + it.effect("keeps malformed tool input non-executable when the provider cleanly ends", () => + Effect.gen(function* () { + const body = sseEvents( + deltaChunk({ + role: "assistant", + tool_calls: [{ index: 0, id: "call_1", function: { name: "lookup", arguments: '{"query":' } }], + }), + ) + const response = yield* LLMClient.generate( + LLM.updateRequest(request, { + tools: [{ name: "lookup", description: "Lookup data", inputSchema: { type: "object" } }], + }), + ).pipe(Effect.provide(fixedResponse(body))) + + expect(response.events.find(LLMEvent.is.toolInputError)).toEqual({ + type: "tool-input-error", + id: "call_1", + name: "lookup", + raw: '{"query":', + }) + expect(response.toolCalls).toEqual([]) + expect(response.finishReason).toBe("unknown") }), )