Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
144 changes: 101 additions & 43 deletions packages/ai/src/protocols/open-responses.ts
Original file line number Diff line number Diff line change
Expand Up @@ -287,6 +287,7 @@ export const Event = Schema.StructWithRest(
delta: Schema.optional(Schema.String),
text: Schema.optional(Schema.String),
item_id: Schema.optional(Schema.String),
output_index: Schema.optional(Schema.Number),
summary_index: Schema.optional(Schema.Number),
item: Schema.optional(StreamItem),
response: Schema.optional(
Expand All @@ -297,6 +298,7 @@ export const Event = Schema.StructWithRest(
incomplete_details: optionalNull(Schema.Struct({ reason: Schema.optional(Schema.String) })),
usage: optionalNull(OpenResponsesUsage),
error: optionalNull(OpenResponsesErrorPayload),
output: Schema.optional(Schema.Array(StreamItem)),
}),
[Schema.Record(Schema.String, Schema.Unknown)],
),
Expand Down Expand Up @@ -342,13 +344,14 @@ export interface ParserState {
readonly messageItems: ReadonlySet<string>
readonly messagePhases: Readonly<Record<string, MessagePhase | null>>
readonly reasoningItems: Readonly<Record<string, ReasoningStreamItem>>
readonly store: boolean | undefined
readonly reasoningIndexes: Readonly<Record<number, string>>
}

type ReasoningSummaryStatus = "active" | "can-conclude" | "concluded"
type ReasoningSummaryStatus = "active" | "concluded"

interface ReasoningStreamItem {
readonly encryptedContent: string | null | undefined
readonly blockIDs?: ReadonlyArray<string>
// Keyed by the wire protocol's numeric `summary_index`. JS object keys coerce to
// strings, but typing the map as `Record<number, ...>` documents intent
// and matches the wire field.
Expand Down Expand Up @@ -857,6 +860,10 @@ const onOutputItemAdded = (state: ParserState, event: Event): StepResult => {
...state.reasoningItems,
[item.id]: { encryptedContent: item.encrypted_content, summaryParts: { 0: "active" } },
},
reasoningIndexes:
event.output_index === undefined
? state.reasoningIndexes
: { ...state.reasoningIndexes, [event.output_index]: item.id },
},
events,
]
Expand Down Expand Up @@ -890,23 +897,11 @@ const onReasoningSummaryPartAdded = (state: ParserState, event: Event): StepResu
if (event.summary_index === 0) return [state, NO_EVENTS]

const events: LLMEvent[] = []
const closed = Object.entries(item.summaryParts)
.filter((entry) => entry[1] === "can-conclude")
.reduce(
(lifecycle, entry) =>
Lifecycle.reasoningEnd(
lifecycle,
events,
`${event.item_id}:${entry[0]}`,
providerMetadata(state, { itemId: event.item_id }),
),
state.lifecycle,
)
return [
{
...state,
lifecycle: Lifecycle.reasoningStart(
closed,
state.lifecycle,
events,
`${event.item_id}:${event.summary_index}`,
providerMetadata(state, { itemId: event.item_id, reasoningEncryptedContent: item.encryptedContent ?? null }),
Expand All @@ -916,11 +911,7 @@ const onReasoningSummaryPartAdded = (state: ParserState, event: Event): StepResu
[event.item_id]: {
...item,
summaryParts: {
...Object.fromEntries(
Object.entries(item.summaryParts).map((entry) =>
entry[1] === "can-conclude" ? [entry[0], "concluded" as const] : entry,
),
),
...item.summaryParts,
[event.summary_index]: "active",
},
},
Expand All @@ -938,22 +929,19 @@ const onReasoningSummaryPartDone = (state: ParserState, event: Event): StepResul
return [
{
...state,
lifecycle:
state.store !== false
? Lifecycle.reasoningEnd(
state.lifecycle,
events,
`${event.item_id}:${event.summary_index}`,
providerMetadata(state, { itemId: event.item_id }),
)
: state.lifecycle,
lifecycle: Lifecycle.reasoningEnd(
state.lifecycle,
events,
`${event.item_id}:${event.summary_index}`,
providerMetadata(state, { itemId: event.item_id }),
),
reasoningItems: {
...state.reasoningItems,
[event.item_id]: {
...item,
summaryParts: {
...item.summaryParts,
[event.summary_index]: state.store !== false ? "concluded" : "can-conclude",
[event.summary_index]: "concluded",
},
},
},
Expand Down Expand Up @@ -1039,27 +1027,74 @@ const onOutputItemDone = Effect.fn("OpenResponses.onOutputItemDone")(function* (
}

if (isReasoningItem(item)) {
if (
event.output_index !== undefined &&
state.reasoningIndexes[event.output_index] !== undefined &&
state.reasoningIndexes[event.output_index] !== item.id
)
return [state, NO_EVENTS] satisfies StepResult
const events: LLMEvent[] = []
const metadata = reasoningMetadata(state, item)
const reasoningItem = state.reasoningItems[item.id]
const reasoningIndexes =
event.output_index === undefined
? state.reasoningIndexes
: Object.fromEntries(
Object.entries(state.reasoningIndexes).filter((entry) => entry[0] !== `${event.output_index}`),
)
if (reasoningItem) {
const lifecycle = Object.entries(reasoningItem.summaryParts)
.filter((entry) => entry[1] === "active" || entry[1] === "can-conclude")
.reduce(
(lifecycle, entry) => Lifecycle.reasoningEnd(lifecycle, events, `${item.id}:${entry[0]}`, metadata),
state.lifecycle,
)
const { [item.id]: _removed, ...reasoningItems } = state.reasoningItems
return [{ ...state, lifecycle, reasoningItems }, events] satisfies StepResult
const openParts = Object.entries(reasoningItem.summaryParts).filter((entry) => entry[1] === "active")
const lifecycle = openParts.reduce(
(lifecycle, entry) => Lifecycle.reasoningEnd(lifecycle, events, `${item.id}:${entry[0]}`, metadata),
state.lifecycle,
)
if (typeof item.encrypted_content === "string" && openParts.length === 0) {
const blockID = Object.keys(reasoningItem.summaryParts)
.map((index) => `${item.id}:${index}`)
.at(-1)
if (blockID) events.push(LLMEvent.reasoningMetadata({ id: blockID, providerMetadata: metadata }))
}
const reasoningItems =
typeof item.encrypted_content === "string"
? Object.fromEntries(Object.entries(state.reasoningItems).filter((entry) => entry[0] !== item.id))
: {
...state.reasoningItems,
[item.id]: {
...reasoningItem,
encryptedContent: item.encrypted_content,
summaryParts: Object.fromEntries(
Object.keys(reasoningItem.summaryParts).map((index) => [index, "concluded" as const]),
),
},
}
return [{ ...state, lifecycle, reasoningItems, reasoningIndexes }, events] satisfies StepResult
}
if (!state.lifecycle.reasoning.has(item.id)) {
const lifecycle = Lifecycle.stepStart(state.lifecycle, events)
events.push(LLMEvent.reasoningStart({ id: item.id, providerMetadata: metadata }))
events.push(LLMEvent.reasoningEnd({ id: item.id, providerMetadata: metadata }))
return [{ ...state, lifecycle }, events] satisfies StepResult
return [
{
...state,
lifecycle,
reasoningIndexes,
reasoningItems:
typeof item.encrypted_content === "string"
? state.reasoningItems
: {
...state.reasoningItems,
[item.id]: { encryptedContent: item.encrypted_content, blockIDs: [item.id], summaryParts: {} },
},
},
events,
] satisfies StepResult
}
return [
{ ...state, lifecycle: Lifecycle.reasoningEnd(state.lifecycle, events, item.id, metadata) },
{
...state,
lifecycle: Lifecycle.reasoningEnd(state.lifecycle, events, item.id, metadata),
reasoningIndexes,
},
events,
] satisfies StepResult
}
Expand All @@ -1077,7 +1112,27 @@ const onResponseFinish = Effect.fn("OpenResponses.onResponseFinish")(function* (
const hasFunctionCall =
pending.events.some((event) => LLMEvent.is.toolCall(event) || LLMEvent.is.toolInputError(event)) ||
state.hasFunctionCall
const lifecycle = Lifecycle.finish(state.lifecycle, events, {
const terminalReasoning = new Map(
(event.response?.output ?? []).filter(isReasoningItem).map((item) => [item.id, item]),
)
const reasoningLifecycle = Object.entries(state.reasoningItems).reduce((lifecycle, [id, item]) => {
const terminal = terminalReasoning.get(id)
if (!terminal) return lifecycle
const metadata = providerMetadata(state, {
itemId: id,
reasoningEncryptedContent: terminal.encrypted_content ?? null,
})
const blockID =
item.blockIDs?.at(-1) ??
Object.keys(item.summaryParts)
.map((index) => `${id}:${index}`)
.at(-1)
if (!blockID) return lifecycle
if (lifecycle.reasoning.has(blockID)) return Lifecycle.reasoningEnd(lifecycle, events, blockID, metadata)
events.push(LLMEvent.reasoningMetadata({ id: blockID, providerMetadata: metadata }))
return lifecycle
}, state.lifecycle)
const lifecycle = Lifecycle.finish(reasoningLifecycle, events, {
reason: {
normalized: mapFinishReason(event, hasFunctionCall),
raw: event.response?.incomplete_details?.reason,
Expand All @@ -1091,7 +1146,10 @@ const onResponseFinish = Effect.fn("OpenResponses.onResponseFinish")(function* (
})
: undefined,
})
return [{ ...state, lifecycle, hasFunctionCall, tools: pending.tools }, events] satisfies StepResult
return [
{ ...state, lifecycle, hasFunctionCall, tools: pending.tools, reasoningItems: {}, reasoningIndexes: {} },
events,
] satisfies StepResult
})

// Build the prettiest summary available from whatever the provider supplied.
Expand Down Expand Up @@ -1210,7 +1268,7 @@ export const initial = (request: LLMRequest, extension: Extension = BASE): Parse
messageItems: new Set<string>(),
messagePhases: {},
reasoningItems: {},
store: OpenResponsesOptions.resolve(request).store,
reasoningIndexes: {},
})

export const protocol = Protocol.make({
Expand Down
25 changes: 25 additions & 0 deletions packages/ai/src/schema/events.ts
Original file line number Diff line number Diff line change
Expand Up @@ -138,6 +138,13 @@ export const ReasoningEnd = Schema.Struct({
}).annotate({ identifier: "LLM.Event.ReasoningEnd" })
export type ReasoningEnd = Schema.Schema.Type<typeof ReasoningEnd>

export const ReasoningMetadata = Schema.Struct({
type: Schema.tag("reasoning-metadata"),
id: ContentBlockID,
providerMetadata: ProviderMetadata,
}).annotate({ identifier: "LLM.Event.ReasoningMetadata" })
export type ReasoningMetadata = Schema.Schema.Type<typeof ReasoningMetadata>

export const ToolInputStart = Schema.Struct({
type: Schema.tag("tool-input-start"),
id: ToolCallID,
Expand Down Expand Up @@ -242,6 +249,7 @@ const llmEventTagged = Schema.Union([
ReasoningStart,
ReasoningDelta,
ReasoningEnd,
ReasoningMetadata,
ToolInputStart,
ToolInputDelta,
ToolInputEnd,
Expand Down Expand Up @@ -278,6 +286,8 @@ export const LLMEvent = Object.assign(llmEventTagged, {
ReasoningDelta.make({ ...input, id: contentBlockID(input.id) }),
reasoningEnd: (input: WithID<ReasoningEnd, ContentBlockID>) =>
ReasoningEnd.make({ ...input, id: contentBlockID(input.id) }),
reasoningMetadata: (input: WithID<ReasoningMetadata, ContentBlockID>) =>
ReasoningMetadata.make({ ...input, id: contentBlockID(input.id) }),
toolInputStart: (input: WithID<ToolInputStart, ToolCallID>) =>
ToolInputStart.make({ ...input, id: toolCallID(input.id) }),
toolInputDelta: (input: WithID<ToolInputDelta, ToolCallID>) =>
Expand Down Expand Up @@ -312,6 +322,7 @@ export const LLMEvent = Object.assign(llmEventTagged, {
reasoningStart: llmEventTagged.guards["reasoning-start"],
reasoningDelta: llmEventTagged.guards["reasoning-delta"],
reasoningEnd: llmEventTagged.guards["reasoning-end"],
reasoningMetadata: llmEventTagged.guards["reasoning-metadata"],
toolInputStart: llmEventTagged.guards["tool-input-start"],
toolInputDelta: llmEventTagged.guards["tool-input-delta"],
toolInputEnd: llmEventTagged.guards["tool-input-end"],
Expand Down Expand Up @@ -483,6 +494,18 @@ const reduceReasoningEnd = (state: ResponseState, event: ReasoningEnd): Response
}
}

const reduceReasoningMetadata = (state: ResponseState, event: ReasoningMetadata): ResponseState => {
const current = state.reasoningParts[event.id]
if (!current) return state
return {
...replaceContent(state, current.contentIndex, reasoningContent(current.text, event.providerMetadata)),
reasoningParts: {
...state.reasoningParts,
[event.id]: { ...current, providerMetadata: event.providerMetadata },
},
}
}

const reduceToolInputStart = (state: ResponseState, event: ToolInputStart): ResponseState => ({
...state,
toolInputs: {
Expand Down Expand Up @@ -552,6 +575,8 @@ const reduceResponseState = (state: ResponseState, event: LLMEvent): ResponseSta
return reduceReasoningDelta(next, event)
case "reasoning-end":
return reduceReasoningEnd(next, event)
case "reasoning-metadata":
return reduceReasoningMetadata(next, event)
case "tool-input-start":
return reduceToolInputStart(next, event)
case "tool-input-delta":
Expand Down
39 changes: 35 additions & 4 deletions packages/ai/test/lib/tool-runtime.ts
Original file line number Diff line number Diff line change
Expand Up @@ -82,16 +82,25 @@ const indexStep = (event: LLMEvent, index: number): LLMEvent => {

const stepState = (events: ReadonlyArray<LLMEvent>) => {
const assistantContent: ContentPart[] = []
const reasoningIndexes = new Map<string, number>()
const toolCalls: ToolCallPart[] = []
let reason: Extract<LLMEvent, { type: "finish" }>["reason"] = { normalized: "unknown" }
let usage: Usage | undefined
let providerMetadata: ProviderMetadata | undefined

for (const event of events) {
if (event.type === "text-delta" || event.type === "reasoning-delta") {
appendText(assistantContent, event.type === "text-delta" ? "text" : "reasoning", event.text)
} else if (event.type === "text-end" || event.type === "reasoning-end") {
appendText(assistantContent, event.type === "text-end" ? "text" : "reasoning", "", event.providerMetadata)
if (event.type === "text-delta") {
appendText(assistantContent, "text", event.text)
} else if (event.type === "reasoning-delta") {
appendReasoning(assistantContent, reasoningIndexes, event.id, event.text, event.providerMetadata)
} else if (event.type === "text-end") {
appendText(assistantContent, "text", "", event.providerMetadata)
} else if (event.type === "reasoning-end") {
appendReasoning(assistantContent, reasoningIndexes, event.id, "", event.providerMetadata)
} else if (event.type === "reasoning-metadata") {
const index = reasoningIndexes.get(event.id)
const reasoning = index === undefined ? undefined : assistantContent[index]
if (reasoning?.type === "reasoning") reasoning.providerMetadata = event.providerMetadata
} else if (event.type === "tool-call") {
assistantContent.push(event)
if (!event.providerExecuted) toolCalls.push(event)
Expand All @@ -114,6 +123,28 @@ const stepState = (events: ReadonlyArray<LLMEvent>) => {
return { assistantContent, toolCalls, reason, usage, providerMetadata }
}

const appendReasoning = (
content: ContentPart[],
indexes: Map<string, number>,
id: string,
text: string,
providerMetadata?: ProviderMetadata,
) => {
const index = indexes.get(id)
if (index === undefined) {
indexes.set(id, content.length)
content.push({ type: "reasoning", text, providerMetadata })
return
}
const current = content[index]
if (current?.type !== "reasoning") return
content[index] = {
...current,
text: `${current.text}${text}`,
providerMetadata: providerMetadata ?? current.providerMetadata,
}
}

const appendText = (
content: ContentPart[],
type: "text" | "reasoning",
Expand Down
Loading
Loading