|
|
@@ -12,6 +12,7 @@ import {
|
|
|
ToolFailure,
|
|
|
ToolResultPart,
|
|
|
type ToolResultValue,
|
|
|
+ Usage,
|
|
|
} from "./schema"
|
|
|
import { type AnyTool, type ExecutableTools, type Tools, toDefinitions } from "./tool"
|
|
|
|
|
|
@@ -72,19 +73,42 @@ export const stream = <T extends Tools>(options: StreamOptions<T>): Stream.Strea
|
|
|
tools: [...options.request.tools.filter((tool) => !runtimeToolNames.has(tool.name)), ...runtimeTools],
|
|
|
})
|
|
|
|
|
|
- const loop = (request: LLMRequest, step: number): Stream.Stream<LLMEvent, LLMError> =>
|
|
|
+ const loop = (
|
|
|
+ request: LLMRequest,
|
|
|
+ step: number,
|
|
|
+ usage: Usage | undefined,
|
|
|
+ providerMetadata: ProviderMetadata | undefined,
|
|
|
+ ): Stream.Stream<LLMEvent, LLMError> =>
|
|
|
Stream.unwrap(
|
|
|
Effect.gen(function* () {
|
|
|
- const state: StepState = { assistantContent: [], toolCalls: [], finishReason: undefined }
|
|
|
+ const state: StepState = {
|
|
|
+ assistantContent: [],
|
|
|
+ toolCalls: [],
|
|
|
+ finishReason: undefined,
|
|
|
+ usage: undefined,
|
|
|
+ providerMetadata: undefined,
|
|
|
+ }
|
|
|
|
|
|
const modelStream = options
|
|
|
.stream(request)
|
|
|
+ .pipe(Stream.map((event) => indexStep(event, step)))
|
|
|
.pipe(Stream.tap((event) => Effect.sync(() => accumulate(state, event))))
|
|
|
+ .pipe(Stream.filter((event) => event.type !== "finish"))
|
|
|
|
|
|
const continuation = Stream.unwrap(
|
|
|
Effect.gen(function* () {
|
|
|
- if (state.finishReason !== "tool-calls" || state.toolCalls.length === 0) return Stream.empty
|
|
|
- if (options.toolExecution === "none") return Stream.empty
|
|
|
+ const totalUsage = addUsage(usage, state.usage)
|
|
|
+ const totalProviderMetadata = mergeProviderMetadata(providerMetadata, state.providerMetadata)
|
|
|
+ const finishStream = Stream.fromIterable([
|
|
|
+ LLMEvent.finish({
|
|
|
+ reason: state.finishReason ?? "unknown",
|
|
|
+ usage: totalUsage,
|
|
|
+ providerMetadata: totalProviderMetadata,
|
|
|
+ }),
|
|
|
+ ])
|
|
|
+
|
|
|
+ if (state.finishReason !== "tool-calls" || state.toolCalls.length === 0) return finishStream
|
|
|
+ if (options.toolExecution === "none") return finishStream
|
|
|
|
|
|
const dispatched = yield* Effect.forEach(
|
|
|
state.toolCalls,
|
|
|
@@ -93,10 +117,14 @@ export const stream = <T extends Tools>(options: StreamOptions<T>): Stream.Strea
|
|
|
)
|
|
|
const resultStream = Stream.fromIterable(dispatched.flatMap(([call, result]) => emitEvents(call, result)))
|
|
|
|
|
|
- if (!options.stopWhen) return resultStream
|
|
|
- if (options.stopWhen({ step, request })) return resultStream
|
|
|
+ if (!options.stopWhen) return resultStream.pipe(Stream.concat(finishStream))
|
|
|
+ if (options.stopWhen({ step, request })) return resultStream.pipe(Stream.concat(finishStream))
|
|
|
|
|
|
- return resultStream.pipe(Stream.concat(loop(followUpRequest(request, state, dispatched), step + 1)))
|
|
|
+ return resultStream.pipe(
|
|
|
+ Stream.concat(
|
|
|
+ loop(followUpRequest(request, state, dispatched), step + 1, totalUsage, totalProviderMetadata),
|
|
|
+ ),
|
|
|
+ )
|
|
|
}),
|
|
|
)
|
|
|
|
|
|
@@ -104,13 +132,21 @@ export const stream = <T extends Tools>(options: StreamOptions<T>): Stream.Strea
|
|
|
}),
|
|
|
)
|
|
|
|
|
|
- return loop(initialRequest, 0)
|
|
|
+ return loop(initialRequest, 0, undefined, undefined)
|
|
|
+}
|
|
|
+
|
|
|
+const indexStep = (event: LLMEvent, index: number): LLMEvent => {
|
|
|
+ if (event.type === "step-start") return LLMEvent.stepStart({ index })
|
|
|
+ if (event.type === "step-finish") return LLMEvent.stepFinish({ ...event, index })
|
|
|
+ return event
|
|
|
}
|
|
|
|
|
|
interface StepState {
|
|
|
assistantContent: ContentPart[]
|
|
|
toolCalls: ToolCallPart[]
|
|
|
finishReason: FinishReason | undefined
|
|
|
+ usage: Usage | undefined
|
|
|
+ providerMetadata: ProviderMetadata | undefined
|
|
|
}
|
|
|
|
|
|
const accumulate = (state: StepState, event: LLMEvent) => {
|
|
|
@@ -154,11 +190,45 @@ const accumulate = (state: StepState, event: LLMEvent) => {
|
|
|
)
|
|
|
return
|
|
|
}
|
|
|
- if (event.type === "step-finish" || event.type === "request-finish") {
|
|
|
+ if (event.type === "step-finish") {
|
|
|
state.finishReason = event.reason === "stop" && state.toolCalls.length > 0 ? "tool-calls" : event.reason
|
|
|
+ state.usage = addUsage(state.usage, event.usage)
|
|
|
+ state.providerMetadata = mergeProviderMetadata(state.providerMetadata, event.providerMetadata)
|
|
|
+ return
|
|
|
+ }
|
|
|
+ if (event.type === "finish") {
|
|
|
+ state.finishReason ??= event.reason
|
|
|
+ state.usage ??= event.usage
|
|
|
+ state.providerMetadata = mergeProviderMetadata(state.providerMetadata, event.providerMetadata)
|
|
|
}
|
|
|
}
|
|
|
|
|
|
+const addUsage = (left: Usage | undefined, right: Usage | undefined) => {
|
|
|
+ if (!left) return right
|
|
|
+ if (!right) return left
|
|
|
+ type UsageKey =
|
|
|
+ | "inputTokens"
|
|
|
+ | "outputTokens"
|
|
|
+ | "nonCachedInputTokens"
|
|
|
+ | "cacheReadInputTokens"
|
|
|
+ | "cacheWriteInputTokens"
|
|
|
+ | "reasoningTokens"
|
|
|
+ | "totalTokens"
|
|
|
+ const sum = (key: UsageKey) =>
|
|
|
+ left[key] === undefined && right[key] === undefined ? undefined : Number(left[key] ?? 0) + Number(right[key] ?? 0)
|
|
|
+
|
|
|
+ return new Usage({
|
|
|
+ inputTokens: sum("inputTokens"),
|
|
|
+ outputTokens: sum("outputTokens"),
|
|
|
+ nonCachedInputTokens: sum("nonCachedInputTokens"),
|
|
|
+ cacheReadInputTokens: sum("cacheReadInputTokens"),
|
|
|
+ cacheWriteInputTokens: sum("cacheWriteInputTokens"),
|
|
|
+ reasoningTokens: sum("reasoningTokens"),
|
|
|
+ totalTokens: sum("totalTokens"),
|
|
|
+ providerMetadata: mergeProviderMetadata(left.providerMetadata, right.providerMetadata),
|
|
|
+ })
|
|
|
+}
|
|
|
+
|
|
|
const sameProviderMetadata = (left: ProviderMetadata | undefined, right: ProviderMetadata | undefined) =>
|
|
|
left === right || JSON.stringify(left) === JSON.stringify(right)
|
|
|
|
|
|
@@ -200,17 +270,17 @@ const dispatch = (tools: Tools, call: ToolCallPart): Effect.Effect<ToolResultVal
|
|
|
if (!tool.execute)
|
|
|
return Effect.succeed({ type: "error" as const, value: `Tool has no execute handler: ${call.name}` })
|
|
|
|
|
|
- return decodeAndExecute(tool, call.input).pipe(
|
|
|
+ return decodeAndExecute(tool, call).pipe(
|
|
|
Effect.catchTag("LLM.ToolFailure", (failure) =>
|
|
|
Effect.succeed({ type: "error" as const, value: failure.message } satisfies ToolResultValue),
|
|
|
),
|
|
|
)
|
|
|
}
|
|
|
|
|
|
-const decodeAndExecute = (tool: AnyTool, input: unknown): Effect.Effect<ToolResultValue, ToolFailure> =>
|
|
|
- tool._decode(input).pipe(
|
|
|
+const decodeAndExecute = (tool: AnyTool, call: ToolCallPart): Effect.Effect<ToolResultValue, ToolFailure> =>
|
|
|
+ tool._decode(call.input).pipe(
|
|
|
Effect.mapError((error) => new ToolFailure({ message: `Invalid tool input: ${error.message}` })),
|
|
|
- Effect.flatMap((decoded) => tool.execute!(decoded)),
|
|
|
+ Effect.flatMap((decoded) => tool.execute!(decoded, { id: call.id, name: call.name })),
|
|
|
Effect.flatMap((value) =>
|
|
|
tool._encode(value).pipe(
|
|
|
Effect.mapError(
|