tool-runtime.ts 5.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151
  1. import { Effect, Stream } from "effect"
  2. import { LLMClient } from "../../src/route"
  3. import {
  4. LLMEvent,
  5. LLMRequest,
  6. Message,
  7. type ContentPart,
  8. type ProviderMetadata,
  9. type ToolCallPart,
  10. ToolResultPart,
  11. type ToolResultValue,
  12. type Usage,
  13. } from "../../src/schema"
  14. import { type Tools, toDefinitions } from "../../src/tool"
  15. import { ToolRuntime } from "../../src/tool-runtime"
  16. interface RunOptions<T extends Tools> {
  17. readonly request: LLMRequest
  18. readonly tools: T
  19. readonly maxSteps?: number
  20. }
  21. /** Test-owned continuation loop. Production callers must own durable history. */
  22. export const runTools = <T extends Tools>(options: RunOptions<T>) =>
  23. Stream.unwrap(
  24. Effect.gen(function* () {
  25. const names = new Set(Object.keys(options.tools))
  26. let request = LLMRequest.update(options.request, {
  27. tools: [...options.request.tools.filter((tool) => !names.has(tool.name)), ...toDefinitions(options.tools)],
  28. })
  29. let usage: Usage | undefined
  30. const events: LLMEvent[] = []
  31. for (let step = 0; step < (options.maxSteps ?? 10); step++) {
  32. const streamed = Array.from(yield* LLMClient.stream(request).pipe(Stream.runCollect))
  33. const state = stepState(streamed)
  34. usage = addUsage(usage, state.usage)
  35. events.push(...streamed.filter((event) => event.type !== "finish").map((event) => indexStep(event, step)))
  36. if (state.toolCalls.length === 0) {
  37. events.push(LLMEvent.finish({ reason: state.reason, usage, providerMetadata: state.providerMetadata }))
  38. return Stream.fromIterable(events)
  39. }
  40. const dispatched = yield* Effect.forEach(
  41. state.toolCalls,
  42. (call) => ToolRuntime.dispatch(options.tools, call).pipe(Effect.map((result) => [call, result] as const)),
  43. { concurrency: 10 },
  44. )
  45. events.push(...dispatched.flatMap(([, result]) => result.events))
  46. if (step + 1 >= (options.maxSteps ?? 10)) {
  47. events.push(LLMEvent.finish({ reason: state.reason, usage, providerMetadata: state.providerMetadata }))
  48. return Stream.fromIterable(events)
  49. }
  50. request = LLMRequest.update(request, {
  51. messages: [
  52. ...request.messages,
  53. Message.assistant(state.assistantContent),
  54. ...dispatched.map(([call, dispatched]) =>
  55. Message.tool({
  56. id: call.id,
  57. name: call.name,
  58. result: dispatched.result,
  59. providerMetadata: call.providerMetadata,
  60. }),
  61. ),
  62. ],
  63. })
  64. }
  65. return Stream.fromIterable(events)
  66. }),
  67. )
  68. const indexStep = (event: LLMEvent, index: number): LLMEvent => {
  69. if (event.type === "step-start") return LLMEvent.stepStart({ index })
  70. if (event.type === "step-finish") return LLMEvent.stepFinish({ ...event, index })
  71. return event
  72. }
  73. const stepState = (events: ReadonlyArray<LLMEvent>) => {
  74. const assistantContent: ContentPart[] = []
  75. const toolCalls: ToolCallPart[] = []
  76. let reason: Extract<LLMEvent, { type: "finish" }>["reason"] = { normalized: "unknown" }
  77. let usage: Usage | undefined
  78. let providerMetadata: ProviderMetadata | undefined
  79. for (const event of events) {
  80. if (event.type === "text-delta" || event.type === "reasoning-delta") {
  81. appendText(assistantContent, event.type === "text-delta" ? "text" : "reasoning", event.text)
  82. } else if (event.type === "text-end" || event.type === "reasoning-end") {
  83. appendText(assistantContent, event.type === "text-end" ? "text" : "reasoning", "", event.providerMetadata)
  84. } else if (event.type === "tool-call") {
  85. assistantContent.push(event)
  86. if (!event.providerExecuted) toolCalls.push(event)
  87. } else if (event.type === "tool-result" && event.providerExecuted && event.result !== undefined) {
  88. assistantContent.push(
  89. ToolResultPart.make({
  90. id: event.id,
  91. name: event.name,
  92. result: event.result,
  93. providerExecuted: true,
  94. providerMetadata: event.providerMetadata,
  95. }),
  96. )
  97. } else if (event.type === "finish") {
  98. reason = event.reason
  99. usage = event.usage
  100. providerMetadata = event.providerMetadata
  101. }
  102. }
  103. return { assistantContent, toolCalls, reason, usage, providerMetadata }
  104. }
  105. const appendText = (
  106. content: ContentPart[],
  107. type: "text" | "reasoning",
  108. text: string,
  109. providerMetadata?: ProviderMetadata,
  110. ) => {
  111. const last = content.at(-1)
  112. if (last?.type === type) {
  113. content[content.length - 1] = {
  114. ...last,
  115. text: `${last.text}${text}`,
  116. providerMetadata: providerMetadata ?? last.providerMetadata,
  117. }
  118. return
  119. }
  120. content.push({ type, text, providerMetadata })
  121. }
  122. const addUsage = (left: Usage | undefined, right: Usage | undefined): Usage | undefined => {
  123. if (!left) return right
  124. if (!right) return left
  125. const sum = (key: keyof Usage) =>
  126. typeof left[key] !== "number" && typeof right[key] !== "number"
  127. ? undefined
  128. : ((left[key] as number | undefined) ?? 0) + ((right[key] as number | undefined) ?? 0)
  129. return {
  130. inputTokens: sum("inputTokens"),
  131. outputTokens: sum("outputTokens"),
  132. nonCachedInputTokens: sum("nonCachedInputTokens"),
  133. cacheReadInputTokens: sum("cacheReadInputTokens"),
  134. cacheWriteInputTokens: sum("cacheWriteInputTokens"),
  135. reasoningTokens: sum("reasoningTokens"),
  136. totalTokens: sum("totalTokens"),
  137. } as Usage
  138. }