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
25 changes: 23 additions & 2 deletions packages/ai/src/protocols/anthropic-messages.ts
Original file line number Diff line number Diff line change
Expand Up @@ -316,7 +316,21 @@ const lowerServerToolResult = Effect.fn("AnthropicMessages.lowerServerToolResult
const wireType = serverToolResultType(part.name)
if (!wireType)
return yield* invalid(`Anthropic Messages does not know how to round-trip server tool result for ${part.name}`)
return { type: wireType, tool_use_id: part.id, content: part.result.value } satisfies AnthropicServerToolResultBlock
const errorType = `${wireType}_error`
const syntheticErrorCode =
ProviderShared.isRecord(part.result.value) &&
ProviderShared.isRecord(part.result.value.error) &&
part.result.value.error.type === "provider.invalid-output"
? "invalid_tool_input"
: "unavailable"
const content =
part.result.type !== "error" ||
(ProviderShared.isRecord(part.result.value) &&
part.result.value.type === errorType &&
typeof part.result.value.error_code === "string")
? part.result.value
: { type: errorType, error_code: syntheticErrorCode }
return { type: wireType, tool_use_id: part.id, content } satisfies AnthropicServerToolResultBlock
})

const lowerImage = Effect.fn("AnthropicMessages.lowerImage")(function* (part: MediaPart) {
Expand Down Expand Up @@ -703,7 +717,14 @@ const onContentBlockStart = (state: ParserState, event: AnthropicEvent): StepRes
providerExecuted: block.type === "server_tool_use",
}),
},
[...events, LLMEvent.toolInputStart({ id: block.id ?? String(event.index), name: block.name ?? "" })],
[
...events,
LLMEvent.toolInputStart({
id: block.id ?? String(event.index),
name: block.name ?? "",
providerExecuted: block.type === "server_tool_use" ? true : undefined,
}),
],
]
}

Expand Down
22 changes: 19 additions & 3 deletions packages/ai/src/protocols/gemini.ts
Original file line number Diff line number Diff line change
Expand Up @@ -441,21 +441,37 @@ const step = (state: ParserState, event: GeminiEvent) => {
if ("functionCall" in part) {
const input = part.functionCall.args
const id = `tool_${nextToolCallId++}`
const providerMetadata = part.thoughtSignature
? googleMetadata({ thoughtSignature: part.thoughtSignature })
: undefined
lifecycle = Lifecycle.reasoningEnd(
lifecycle,
events,
"reasoning-0",
reasoningSignature ? googleMetadata({ thoughtSignature: reasoningSignature }) : undefined,
)
lifecycle = Lifecycle.stepStart(lifecycle, events)
if (typeof input === "string") {
events.push(
LLMEvent.toolInputStart({ id, name: part.functionCall.name, providerMetadata }),
LLMEvent.toolInputEnd({ id, name: part.functionCall.name, input, providerMetadata }),
LLMEvent.toolInputError({
id,
name: part.functionCall.name,
raw: input,
message: `Invalid JSON input for ${ADAPTER} tool call ${part.functionCall.name}`,
providerMetadata,
}),
)
hasToolCalls = true
continue
}
events.push(
LLMEvent.toolCall({
id,
name: part.functionCall.name,
input,
providerMetadata: part.thoughtSignature
? googleMetadata({ thoughtSignature: part.thoughtSignature })
: undefined,
providerMetadata,
}),
)
hasToolCalls = true
Expand Down
19 changes: 15 additions & 4 deletions packages/ai/src/protocols/openai-responses.ts
Original file line number Diff line number Diff line change
Expand Up @@ -820,22 +820,33 @@ const onOutputItemDone = Effect.fn("OpenAIResponses.onOutputItemDone")(function*

if (item.type === "function_call") {
if (!item.id || !item.call_id || !item.name) return [state, NO_EVENTS] satisfies StepResult
const tools = state.tools[item.id]
const existing = state.tools[item.id]
const providerMetadata = openaiMetadata({ itemId: item.id })
const tools = existing
? state.tools
: ToolStream.start(state.tools, item.id, { id: item.call_id, name: item.name })
: ToolStream.start(state.tools, item.id, {
id: item.call_id,
name: item.name,
providerMetadata,
})
const result =
item.arguments === undefined
? yield* ToolStream.finish(ADAPTER, tools, item.id)
: yield* ToolStream.finishWithInput(ADAPTER, tools, item.id, item.arguments)
const events: LLMEvent[] = []
const resultEvents = result.events ?? []
const resultEvents = [
...(existing ? [] : [LLMEvent.toolInputStart({ id: item.call_id, name: item.name, providerMetadata })]),
...(result.events ?? []),
]
const lifecycle = resultEvents.length ? Lifecycle.stepStart(state.lifecycle, events) : state.lifecycle
events.push(...resultEvents)
return [
{
...state,
lifecycle,
hasFunctionCall: resultEvents.some(LLMEvent.is.toolCall) ? true : state.hasFunctionCall,
hasFunctionCall:
state.hasFunctionCall ||
resultEvents.some((event) => LLMEvent.is.toolCall(event) || LLMEvent.is.toolInputError(event)),
tools: result.tools,
},
events,
Expand Down
31 changes: 29 additions & 2 deletions packages/ai/src/protocols/shared.ts
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import {
type ContentPart,
type LLMRequest,
type MediaPart,
type ProviderMetadata,
type ToolFileContent,
type TextPart,
type ToolResultPart,
Expand Down Expand Up @@ -152,8 +153,34 @@ export const wrappedSystemUpdate = Effect.fn("ProviderShared.wrappedSystemUpdate
* input deltas (e.g. zero-arg tools). The error message is uniform across
* routes: `Invalid JSON input for <route> tool call <name>`.
*/
export const parseToolInput = (route: string, name: string, raw: string) =>
parseJson(route, raw || "{}", `Invalid JSON input for ${route} tool call ${name}`)
export const parseToolInput = (
route: string,
tool: {
readonly id: string
readonly name: string
readonly providerExecuted?: boolean
readonly providerMetadata?: ProviderMetadata
},
raw: string,
) =>
Effect.try({
try: () => decodeJson(raw || "{}"),
catch: () =>
new LLMError({
module: "ProviderShared",
method: "stream",
reason: new InvalidProviderOutputReason({
route,
message: `Invalid JSON input for ${route} tool call ${tool.name}`,
raw,
source: "tool-input",
toolCallID: tool.id,
toolName: tool.name,
providerExecuted: tool.providerExecuted,
providerMetadata: tool.providerMetadata,
}),
}),
})

export const IMAGE_MIMES = ["image/png", "image/jpeg", "image/gif", "image/webp"] as const
export const VIDEO_MIMES = ["video/mp4", "video/webm", "video/quicktime"] as const
Expand Down
30 changes: 25 additions & 5 deletions packages/ai/src/protocols/utils/tool-stream.ts
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@ const inputStart = (tool: PendingTool) =>
LLMEvent.toolInputStart({
id: tool.id,
name: tool.name,
providerExecuted: tool.providerExecuted ? true : undefined,
providerMetadata: tool.providerMetadata,
})

Expand All @@ -63,8 +64,9 @@ const inputDelta = (tool: PendingTool, text: string) =>
text,
})

const toolCall = (route: string, tool: PendingTool, inputOverride?: string) =>
parseToolInput(route, tool.name, inputOverride ?? tool.input).pipe(
const toolCall = (route: string, tool: PendingTool, inputOverride?: string) => {
const raw = inputOverride ?? tool.input
return parseToolInput(route, tool, raw).pipe(
Effect.map(
(input): ToolCall =>
LLMEvent.toolCall({
Expand All @@ -75,7 +77,20 @@ const toolCall = (route: string, tool: PendingTool, inputOverride?: string) =>
providerMetadata: tool.providerMetadata,
}),
),
Effect.match({
onFailure: (error) =>
LLMEvent.toolInputError({
id: tool.id,
name: tool.name,
raw,
message: error.reason.message,
providerExecuted: tool.providerExecuted ? true : undefined,
providerMetadata: tool.providerMetadata,
}),
onSuccess: (event) => event,
}),
)
}

/** Store the updated tool and produce the optional public delta event. */
const appendTool = <K extends StreamKey>(
Expand Down Expand Up @@ -158,8 +173,8 @@ export const appendExisting = <K extends StreamKey>(

/**
* Finalize one pending tool call: parse the accumulated raw JSON, remove it
* from state, and return the optional public `tool-call` event. Missing keys are
* a no-op because some providers emit stop events for non-tool content blocks.
* from state, and emit either `tool-call` or `tool-input-error`. Missing keys
* are a no-op because some providers emit stop events for non-tool blocks.
*/
export const finish = <K extends StreamKey>(route: string, tools: State<K>, key: K) =>
Effect.gen(function* () {
Expand All @@ -186,7 +201,12 @@ export const finishWithInput = <K extends StreamKey>(route: string, tools: State
return {
tools: withoutTool(tools, key),
events: [
LLMEvent.toolInputEnd({ id: tool.id, name: tool.name, providerMetadata: tool.providerMetadata }),
LLMEvent.toolInputEnd({
id: tool.id,
name: tool.name,
input,
providerMetadata: tool.providerMetadata,
}),
yield* toolCall(route, tool, input),
],
}
Expand Down
4 changes: 4 additions & 0 deletions packages/ai/src/schema/errors.ts
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,10 @@ export class InvalidProviderOutputReason extends Schema.Class<InvalidProviderOut
message: Schema.String,
route: Schema.optional(Schema.String),
raw: Schema.optional(Schema.String),
source: Schema.optional(Schema.Literal("tool-input")),
toolCallID: Schema.optional(Schema.String),
toolName: Schema.optional(Schema.String),
providerExecuted: Schema.optional(Schema.Boolean),
providerMetadata: Schema.optional(ProviderMetadata),
}) {}

Expand Down
20 changes: 20 additions & 0 deletions packages/ai/src/schema/events.ts
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,7 @@ export const ToolInputStart = Schema.Struct({
type: Schema.tag("tool-input-start"),
id: ToolCallID,
name: Schema.String,
providerExecuted: Schema.optional(Schema.Boolean),
providerMetadata: Schema.optional(ProviderMetadata),
}).annotate({ identifier: "LLM.Event.ToolInputStart" })
export type ToolInputStart = Schema.Schema.Type<typeof ToolInputStart>
Expand All @@ -145,10 +146,22 @@ export const ToolInputEnd = Schema.Struct({
type: Schema.tag("tool-input-end"),
id: ToolCallID,
name: Schema.String,
input: Schema.optional(Schema.String),
providerMetadata: Schema.optional(ProviderMetadata),
}).annotate({ identifier: "LLM.Event.ToolInputEnd" })
export type ToolInputEnd = Schema.Schema.Type<typeof ToolInputEnd>

export const ToolInputError = Schema.Struct({
type: Schema.tag("tool-input-error"),
id: ToolCallID,
name: Schema.String,
raw: Schema.String,
message: Schema.String,
providerExecuted: Schema.optional(Schema.Boolean),
providerMetadata: Schema.optional(ProviderMetadata),
}).annotate({ identifier: "LLM.Event.ToolInputError" })
export type ToolInputError = Schema.Schema.Type<typeof ToolInputError>

export const ToolCall = Schema.Struct({
type: Schema.tag("tool-call"),
id: ToolCallID,
Expand Down Expand Up @@ -216,6 +229,7 @@ const llmEventTagged = Schema.Union([
ToolInputStart,
ToolInputDelta,
ToolInputEnd,
ToolInputError,
ToolCall,
ToolResult,
ToolError,
Expand Down Expand Up @@ -253,6 +267,8 @@ export const LLMEvent = Object.assign(llmEventTagged, {
toolInputDelta: (input: WithID<ToolInputDelta, ToolCallID>) =>
ToolInputDelta.make({ ...input, id: toolCallID(input.id) }),
toolInputEnd: (input: WithID<ToolInputEnd, ToolCallID>) => ToolInputEnd.make({ ...input, id: toolCallID(input.id) }),
toolInputError: (input: WithID<ToolInputError, ToolCallID>) =>
ToolInputError.make({ ...input, id: toolCallID(input.id) }),
toolCall: (input: WithID<ToolCall, ToolCallID>) => ToolCall.make({ ...input, id: toolCallID(input.id) }),
toolResult: (input: WithID<ToolResult, ToolCallID>) =>
ToolResult.make({
Expand Down Expand Up @@ -283,6 +299,7 @@ export const LLMEvent = Object.assign(llmEventTagged, {
toolInputStart: llmEventTagged.guards["tool-input-start"],
toolInputDelta: llmEventTagged.guards["tool-input-delta"],
toolInputEnd: llmEventTagged.guards["tool-input-end"],
toolInputError: llmEventTagged.guards["tool-input-error"],
toolCall: llmEventTagged.guards["tool-call"],
toolResult: llmEventTagged.guards["tool-result"],
toolError: llmEventTagged.guards["tool-error"],
Expand Down Expand Up @@ -498,6 +515,7 @@ const reduceToolInputEnd = (state: ResponseState, event: ToolInputEnd): Response
[event.id]: {
...current,
name: event.name,
text: event.input ?? current.text,
providerMetadata: event.providerMetadata ?? current.providerMetadata,
},
},
Expand Down Expand Up @@ -548,6 +566,8 @@ const reduceResponseState = (state: ResponseState, event: LLMEvent): ResponseSta
return reduceToolInputDelta(next, event)
case "tool-input-end":
return reduceToolInputEnd(next, event)
case "tool-input-error":
return next
case "tool-call":
return reduceToolCall(next, event)
case "tool-result":
Expand Down
Loading
Loading