tool-runtime.ts 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146
  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({ id: call.id, name: call.name, result: dispatched.result }),
  56. ),
  57. ],
  58. })
  59. }
  60. return Stream.fromIterable(events)
  61. }),
  62. )
  63. const indexStep = (event: LLMEvent, index: number): LLMEvent => {
  64. if (event.type === "step-start") return LLMEvent.stepStart({ index })
  65. if (event.type === "step-finish") return LLMEvent.stepFinish({ ...event, index })
  66. return event
  67. }
  68. const stepState = (events: ReadonlyArray<LLMEvent>) => {
  69. const assistantContent: ContentPart[] = []
  70. const toolCalls: ToolCallPart[] = []
  71. let reason: Extract<LLMEvent, { type: "finish" }>["reason"] = "unknown"
  72. let usage: Usage | undefined
  73. let providerMetadata: ProviderMetadata | undefined
  74. for (const event of events) {
  75. if (event.type === "text-delta" || event.type === "reasoning-delta") {
  76. appendText(assistantContent, event.type === "text-delta" ? "text" : "reasoning", event.text)
  77. } else if (event.type === "text-end" || event.type === "reasoning-end") {
  78. appendText(assistantContent, event.type === "text-end" ? "text" : "reasoning", "", event.providerMetadata)
  79. } else if (event.type === "tool-call") {
  80. assistantContent.push(event)
  81. if (!event.providerExecuted) toolCalls.push(event)
  82. } else if (event.type === "tool-result" && event.providerExecuted && event.result !== undefined) {
  83. assistantContent.push(
  84. ToolResultPart.make({
  85. id: event.id,
  86. name: event.name,
  87. result: event.result,
  88. providerExecuted: true,
  89. providerMetadata: event.providerMetadata,
  90. }),
  91. )
  92. } else if (event.type === "finish") {
  93. reason = event.reason
  94. usage = event.usage
  95. providerMetadata = event.providerMetadata
  96. }
  97. }
  98. return { assistantContent, toolCalls, reason, usage, providerMetadata }
  99. }
  100. const appendText = (
  101. content: ContentPart[],
  102. type: "text" | "reasoning",
  103. text: string,
  104. providerMetadata?: ProviderMetadata,
  105. ) => {
  106. const last = content.at(-1)
  107. if (last?.type === type) {
  108. content[content.length - 1] = {
  109. ...last,
  110. text: `${last.text}${text}`,
  111. providerMetadata: providerMetadata ?? last.providerMetadata,
  112. }
  113. return
  114. }
  115. content.push({ type, text, providerMetadata })
  116. }
  117. const addUsage = (left: Usage | undefined, right: Usage | undefined): Usage | undefined => {
  118. if (!left) return right
  119. if (!right) return left
  120. const sum = (key: keyof Usage) =>
  121. typeof left[key] !== "number" && typeof right[key] !== "number"
  122. ? undefined
  123. : ((left[key] as number | undefined) ?? 0) + ((right[key] as number | undefined) ?? 0)
  124. return {
  125. inputTokens: sum("inputTokens"),
  126. outputTokens: sum("outputTokens"),
  127. nonCachedInputTokens: sum("nonCachedInputTokens"),
  128. cacheReadInputTokens: sum("cacheReadInputTokens"),
  129. cacheWriteInputTokens: sum("cacheWriteInputTokens"),
  130. reasoningTokens: sum("reasoningTokens"),
  131. totalTokens: sum("totalTokens"),
  132. } as Usage
  133. }