From 6bb520046469fd1453c1f7cc3ea501dd250dfbaf Mon Sep 17 00:00:00 2001 From: Aiden Cline <63023139+rekram1-node@users.noreply.github.com> Date: Mon, 24 Aug 2026 22:10:21 -0500 Subject: [PATCH] fix(ai): enforce chat finish reasons (#44743) --- packages/ai/src/protocols/bedrock-converse.ts | 2 +- packages/ai/src/protocols/gemini.ts | 2 +- packages/ai/src/protocols/openai-chat.ts | 109 +++++++++++++----- packages/ai/src/route/client.ts | 30 ++++- packages/ai/src/route/protocol.ts | 4 +- .../ai/test/provider/bedrock-mantle.test.ts | 5 +- packages/ai/test/provider/openai-chat.test.ts | 6 + .../provider/openai-compatible-chat.test.ts | 103 ++++++++++++++++- 8 files changed, 220 insertions(+), 41 deletions(-) diff --git a/packages/ai/src/protocols/bedrock-converse.ts b/packages/ai/src/protocols/bedrock-converse.ts index efe3a3ba5dc..3740f10f89d 100644 --- a/packages/ai/src/protocols/bedrock-converse.ts +++ b/packages/ai/src/protocols/bedrock-converse.ts @@ -716,7 +716,7 @@ export const protocol = Protocol.make({ reasoningSignatures: {}, }), step, - onHalt, + onHalt: (state) => Effect.succeed(onHalt(state)), }, }) diff --git a/packages/ai/src/protocols/gemini.ts b/packages/ai/src/protocols/gemini.ts index 21106a91446..7cf59d7eb89 100644 --- a/packages/ai/src/protocols/gemini.ts +++ b/packages/ai/src/protocols/gemini.ts @@ -718,7 +718,7 @@ export const protocol = Protocol.make({ lifecycle: Lifecycle.initial(), }), step, - onHalt: finish, + onHalt: (state) => Effect.succeed(finish(state)), }, }) diff --git a/packages/ai/src/protocols/openai-chat.ts b/packages/ai/src/protocols/openai-chat.ts index fbfdcf64f09..2821db8cf87 100644 --- a/packages/ai/src/protocols/openai-chat.ts +++ b/packages/ai/src/protocols/openai-chat.ts @@ -8,6 +8,8 @@ import { Protocol } from "../route/protocol.js" import { AIError, LLMEvent, + ProviderInternalReason, + UnknownProviderReason, Usage, type FinishReason, type FinishReasonDetails, @@ -224,16 +226,22 @@ const OpenAIChatChoice = Schema.StructWithRest( [Schema.Record(Schema.String, Schema.Unknown)], ) -const OpenAIChatError = Schema.Struct({ - code: optionalNull(Schema.Union([Schema.String, Schema.Number])), - message: Schema.String, -}) +const OpenAIChatError = Schema.StructWithRest( + Schema.Struct({ + code: optionalNull(Schema.Union([Schema.String, Schema.Number])), + message: Schema.String, + }), + [Schema.Record(Schema.String, Schema.Unknown)], +) -export const OpenAIChatEvent = Schema.Struct({ - choices: optionalNull(Schema.Array(OpenAIChatChoice)), - usage: optionalNull(OpenAIChatUsage), - error: optionalNull(OpenAIChatError), -}) +export const OpenAIChatEvent = Schema.StructWithRest( + Schema.Struct({ + choices: optionalNull(Schema.Array(OpenAIChatChoice)), + usage: optionalNull(OpenAIChatUsage), + error: optionalNull(OpenAIChatError), + }), + [Schema.Record(Schema.String, Schema.Unknown)], +) export type OpenAIChatEvent = Schema.Schema.Type type OpenAIChatRequestMessage = LLMRequest["messages"][number] @@ -256,6 +264,7 @@ export interface ParserState { readonly reasoningEmitted: boolean readonly latestToolIndex?: number readonly nextToolIndex: number + readonly requireFinishReason: boolean } // ============================================================================= @@ -726,14 +735,40 @@ export const fromRequest = Effect.fn("OpenAIChat.fromRequest")(function* ( // Streaming parsers are small state machines: every event returns a new state // plus the common `LLMEvent`s produced by that event. Tool calls are accumulated // because OpenAI streams JSON arguments across multiple deltas. -const mapFinishReason = (reason: string | null | undefined): FinishReason => { - if (reason === "stop") return "stop" - if (reason === "length") return "length" - if (reason === "content_filter") return "content-filter" - if (reason === "function_call" || reason === "tool_calls") return "tool-calls" - if (reason === "error") return "error" - return "unknown" -} +const finishReasonError = (event: OpenAIChatEvent, reason: AIError["reason"]) => + new AIError({ + module: ADAPTER, + method: "stream", + body: ProviderShared.encodeJson(event), + reason, + }) + +const mapFinishReason = Effect.fn("OpenAIChat.mapFinishReason")(function* (event: OpenAIChatEvent, reason: string) { + switch (reason) { + case "error": + return yield* finishReasonError( + event, + new UnknownProviderReason({ message: "Provider reported an error (finish_reason: error)" }), + ) + case "network_error": + return yield* finishReasonError( + event, + new ProviderInternalReason({ message: "Provider reported a network error (finish_reason: network_error)" }), + ) + case "stop": + case "end": + return "stop" as const + case "length": + return "length" as const + case "content_filter": + return "content-filter" as const + case "function_call": + case "tool_calls": + return "tool-calls" as const + default: + return "unknown" as const + } +}) // OpenAI Chat reports `prompt_tokens` (inclusive total) with a // cached-read and cache-write subsets, and `completion_tokens` (inclusive @@ -846,16 +881,20 @@ const reasoningMetadata = (field: ParserState["reasoningField"], details?: Reado const step = (state: ParserState, event: OpenAIChatEvent) => Effect.gen(function* () { - if (event.error) + if (event.error) { + const body = ProviderShared.encodeJson(event) return yield* new AIError({ module: ADAPTER, method: "stream", + body, reason: classifyProviderFailure({ message: event.error.message, code: event.error.code === undefined || event.error.code === null ? undefined : String(event.error.code), status: typeof event.error.code === "number" ? event.error.code : undefined, + rawBody: body, }), }) + } const events: LLMEvent[] = [] const choice = event.choices?.[0] // Moonshot (and a few other OpenAI-compatible providers) attach usage to @@ -864,8 +903,11 @@ const step = (state: ParserState, event: OpenAIChatEvent) => const usage = mapUsage(event.usage) ?? (choiceUsage ? mapUsage(choiceUsage) : undefined) ?? state.usage const rawFinishReason = choice?.finish_reason const finishReason = - rawFinishReason !== undefined && rawFinishReason !== null - ? { normalized: mapFinishReason(rawFinishReason), raw: choice?.native_finish_reason ?? rawFinishReason } + rawFinishReason + ? { + normalized: yield* mapFinishReason(event, rawFinishReason), + raw: choice?.native_finish_reason ?? rawFinishReason, + } : state.finishReason const delta = choice?.delta const toolDeltas = delta?.tool_calls ?? [] @@ -885,7 +927,11 @@ const step = (state: ParserState, event: OpenAIChatEvent) => toolDeltas.some((tool) => Boolean(tool.id) || Boolean(tool.function?.name) || Boolean(tool.function?.arguments)) if (state.finishReason !== undefined) { if (hasLateContent) - return yield* ProviderShared.eventError(ADAPTER, "OpenAI Chat received content after the finish reason") + return yield* ProviderShared.eventError( + ADAPTER, + "OpenAI Chat received content after the finish reason", + ProviderShared.encodeJson(event), + ) return [{ ...state, usage }, events] as const } @@ -957,14 +1003,19 @@ const step = (state: ParserState, event: OpenAIChatEvent) => { id: id || undefined, name: name || undefined, text }, "OpenAI Chat tool call delta is missing id or name", ) - if (ToolStream.isError(result)) return yield* result + if (ToolStream.isError(result)) + return yield* ProviderShared.eventError(ADAPTER, result.reason.message, ProviderShared.encodeJson(event)) tools = result.tools if (result.events.length) lifecycle = Lifecycle.stepStart(lifecycle, events) events.push(...result.events) } if (finishReason !== undefined && state.finishReason === undefined && Object.keys(pendingTools).length > 0) - return yield* ProviderShared.eventError(ADAPTER, "OpenAI Chat tool call delta is missing id or name") + return yield* ProviderShared.eventError( + ADAPTER, + "OpenAI Chat tool call delta is missing id or name", + ProviderShared.encodeJson(event), + ) // Finalize accumulated tool inputs eagerly when finish_reason arrives so // valid calls and malformed local calls settle independently. @@ -987,16 +1038,19 @@ const step = (state: ParserState, event: OpenAIChatEvent) => reasoningEmitted, latestToolIndex, nextToolIndex, + requireFinishReason: state.requireFinishReason, }, events, ] as const }) -const finishEvents = (state: ParserState): ReadonlyArray => { +const finishEvents = Effect.fn("OpenAIChat.finishEvents")(function* (state: ParserState) { + if (state.finishReason === undefined && state.requireFinishReason) + return yield* ProviderShared.eventError(ADAPTER, "OpenAI Chat stream ended without finish_reason") const events: LLMEvent[] = [] const toolCallEvents = state.finishReason === undefined && Object.keys(state.tools).length > 0 - ? Effect.runSync(ToolStream.finishAll(ADAPTER, state.tools)).events + ? (yield* ToolStream.finishAll(ADAPTER, state.tools)).events : state.toolCallEvents const hasToolCalls = toolCallEvents.length > 0 const reason = state.finishReason @@ -1005,7 +1059,7 @@ const finishEvents = (state: ParserState): ReadonlyArray => { normalized: state.finishReason.normalized === "stop" && hasToolCalls ? "tool-calls" : state.finishReason.normalized, } - : { normalized: hasToolCalls ? ("tool-calls" as const) : ("unknown" as const) } + : { normalized: hasToolCalls ? ("tool-calls" as const) : ("stop" as const) } const metadata = reasoningMetadata( state.reasoningField, state.reasoningDetailsObserved ? state.reasoningDetails : undefined, @@ -1019,7 +1073,7 @@ const finishEvents = (state: ParserState): ReadonlyArray => { events.push(...toolCallEvents) Lifecycle.finish(lifecycle, events, { reason, usage: state.usage }) return events -} +}) // ============================================================================= // Protocol And OpenAI Route @@ -1048,6 +1102,7 @@ export const protocol = Protocol.make({ reasoningDetailsObserved: false, reasoningEmitted: false, nextToolIndex: 0, + requireFinishReason: request.model.compatibility?.requireFinishReason ?? true, }), step, onHalt: finishEvents, diff --git a/packages/ai/src/route/client.ts b/packages/ai/src/route/client.ts index e7e5d10bf77..0c3090a2716 100644 --- a/packages/ai/src/route/client.ts +++ b/packages/ai/src/route/client.ts @@ -321,12 +321,30 @@ function makeFromTransport( Stream.mapEffect(decodeEvent(route)), protocol.stream.terminal ? Stream.takeUntil(protocol.stream.terminal) : (stream) => stream, ) - const stream = events.pipe( - Stream.mapAccumEffect( - () => protocol.stream.initial(request), - protocol.stream.step, - protocol.stream.onHalt ? { onHalt: protocol.stream.onHalt } : undefined, - ), + const 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 + return onHalt + ? parsed.pipe( + Stream.concat( + Stream.suspend(() => + Stream.unwrap(onHalt(state).pipe(Effect.map(Stream.fromIterable))), + ), + ), + ) + : parsed + }).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 83320440e08..f3a844ec351 100644 --- a/packages/ai/src/route/protocol.ts +++ b/packages/ai/src/route/protocol.ts @@ -59,8 +59,8 @@ export interface ProtocolStream { readonly step: (state: State, event: Event) => Effect.Effect], AIError> /** 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, AIError> } /** diff --git a/packages/ai/test/provider/bedrock-mantle.test.ts b/packages/ai/test/provider/bedrock-mantle.test.ts index f39f0567f92..a56ef940415 100644 --- a/packages/ai/test/provider/bedrock-mantle.test.ts +++ b/packages/ai/test/provider/bedrock-mantle.test.ts @@ -6,6 +6,7 @@ import { AmazonBedrockMantle } from "../../src/providers.js" import { compileRequest, LLMClient } from "../../src/route/client.js" import { it } from "../lib/effect.js" import { dynamicResponse } from "../lib/http.js" +import { sseEvents } from "../lib/sse.js" import { recordedTests } from "../recorded-test.js" const credentials = { @@ -71,7 +72,9 @@ describe("Amazon Bedrock Mantle provider", () => { Effect.gen(function* () { const request = yield* HttpClientRequest.toWeb(input.request) seen.push({ url: request.url, authorization: request.headers.get("authorization") ?? undefined }) - return input.respond("", { headers: { "content-type": "text/event-stream" } }) + return input.respond(sseEvents({ choices: [{ delta: {}, finish_reason: "stop" }] }), { + headers: { "content-type": "text/event-stream" }, + }) }), ), ), diff --git a/packages/ai/test/provider/openai-chat.test.ts b/packages/ai/test/provider/openai-chat.test.ts index 4e9104ad95a..b681559126d 100644 --- a/packages/ai/test/provider/openai-chat.test.ts +++ b/packages/ai/test/provider/openai-chat.test.ts @@ -1251,6 +1251,11 @@ describe("OpenAI Chat route", () => { ).pipe(Effect.provide(fixedResponse(body)), Effect.flip) expect(error.message).toContain("OpenAI Chat tool call delta is missing id or name") + expect(error.reason._tag).toBe("InvalidProviderOutput") + if (error.reason._tag !== "InvalidProviderOutput") return + expect(decodeJson(error.reason.raw ?? "")).toMatchObject({ + choices: [{ finish_reason: "tool_calls" }], + }) }), ) @@ -1264,6 +1269,7 @@ describe("OpenAI Chat route", () => { deltaChunk({ tool_calls: [{ index: 0, function: { arguments: ':"weather"}' } }] }), ) const input = LLMRequest.update(request, { + model: LanguageModel.update(model, { compatibility: { requireFinishReason: false } }), tools: [ToolDefinition.make({ name: "lookup", description: "Lookup data", inputSchema: { type: "object" } })], }) const response = yield* LLMClient.generate(input).pipe(Effect.provide(fixedResponse(body))) diff --git a/packages/ai/test/provider/openai-compatible-chat.test.ts b/packages/ai/test/provider/openai-compatible-chat.test.ts index a7f8ed84575..01144c7e26c 100644 --- a/packages/ai/test/provider/openai-compatible-chat.test.ts +++ b/packages/ai/test/provider/openai-compatible-chat.test.ts @@ -353,13 +353,105 @@ describe("OpenAI-compatible Chat route", () => { }), ) - it.effect("treats an empty finish reason as terminal", () => + it.effect("rejects a stream without a required finish reason", () => Effect.gen(function* () { - const response = yield* LLMClient.generate(request).pipe( + const error = yield* LLMClient.generate(request).pipe( + Effect.provide(fixedResponse(sseEvents(deltaChunk({ content: "Hello" }), deltaChunk({}, "")))), + Effect.flip, + ) + + expect(error.reason).toMatchObject({ + _tag: "InvalidProviderOutput", + message: "OpenAI Chat stream ended without finish_reason", + }) + }), + ) + + it.effect("infers stop when finish reasons are optional", () => + Effect.gen(function* () { + const compatible = OpenAICompatibleChat.route + .with({ provider: "custom", endpoint: { baseURL: "https://api.custom.test/v1" } }) + .model({ id: "custom-model", compatibility: { requireFinishReason: false } }) + const response = yield* LLMClient.generate(LLMRequest.update(request, { model: compatible })).pipe( Effect.provide(fixedResponse(sseEvents(deltaChunk({ content: "Hello" }), deltaChunk({}, "")))), ) - expect(response.finishReason).toEqual({ normalized: "unknown", raw: "" }) + expect(response.finishReason).toEqual({ normalized: "stop" }) + }), + ) + + it.effect("normalizes the end finish reason to stop", () => + Effect.gen(function* () { + const response = yield* LLMClient.generate(request).pipe( + Effect.provide(fixedResponse(sseEvents(deltaChunk({ content: "Hello" }), deltaChunk({}, "end")))), + ) + + expect(response.finishReason).toEqual({ normalized: "stop", raw: "end" }) + }), + ) + + it.effect("classifies provider error finish reasons", () => + Effect.gen(function* () { + const error = yield* LLMClient.generate(request).pipe( + Effect.provide(fixedResponse(sseEvents(deltaChunk({}, "network_error")))), + Effect.flip, + ) + + expect(error.reason).toMatchObject({ + _tag: "ProviderInternal", + message: "Provider reported a network error (finish_reason: network_error)", + }) + expect(decodeJson(error.body ?? "")).toMatchObject({ + id: "chatcmpl_fixture", + choices: [{ finish_reason: "network_error" }], + }) + + const generic = yield* LLMClient.generate(request).pipe( + Effect.provide(fixedResponse(sseEvents(deltaChunk({}, "error")))), + Effect.flip, + ) + expect(generic.reason).toMatchObject({ + _tag: "UnknownProvider", + message: "Provider reported an error (finish_reason: error)", + }) + }), + ) + + it.effect("preserves explicit provider error events", () => + Effect.gen(function* () { + const error = yield* LLMClient.generate(request).pipe( + Effect.provide( + fixedResponse( + sseEvents({ + id: "chatcmpl_error", + error: { code: 502, message: "Provider disconnected", details: { upstream: "vendor" } }, + trace_id: "trace_1", + }), + ), + ), + Effect.flip, + ) + + expect(error.reason).toMatchObject({ _tag: "ProviderInternal", message: "Provider disconnected", status: 502 }) + expect(decodeJson(error.body ?? "")).toMatchObject({ + id: "chatcmpl_error", + error: { code: 502, message: "Provider disconnected", details: { upstream: "vendor" } }, + trace_id: "trace_1", + }) + }), + ) + + it.effect("preserves provider finish outcomes in the common reason algebra", () => + Effect.gen(function* () { + const filtered = yield* LLMClient.generate(request).pipe( + Effect.provide(fixedResponse(sseEvents(deltaChunk({}, "content_filter")))), + ) + const future = yield* LLMClient.generate(request).pipe( + Effect.provide(fixedResponse(sseEvents(deltaChunk({}, "future_reason")))), + ) + + expect(filtered.finishReason).toEqual({ normalized: "content-filter", raw: "content_filter" }) + expect(future.finishReason).toEqual({ normalized: "unknown", raw: "future_reason" }) }), ) @@ -379,6 +471,11 @@ describe("OpenAI-compatible Chat route", () => { ) expect(error.message).toContain("OpenAI Chat received content after the finish reason") + expect(error.reason._tag).toBe("InvalidProviderOutput") + if (error.reason._tag !== "InvalidProviderOutput") return + expect(decodeJson(error.reason.raw ?? "")).toMatchObject({ + choices: [{ delta: { tool_calls: [{ id: "call_1" }] } }], + }) }), ) })