File size: 5,327 Bytes
eb1202c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
import { Effect, Stream } from "effect"
import { LLMClient } from "../../src/route"
import {
  LLMEvent,
  LLMRequest,
  Message,
  type ContentPart,
  type ProviderMetadata,
  type ToolCallPart,
  ToolResultPart,
  type ToolResultValue,
  type Usage,
} from "../../src/schema"
import { type Tools, toDefinitions } from "../../src/tool"
import { ToolRuntime } from "../../src/tool-runtime"

interface RunOptions<T extends Tools> {
  readonly request: LLMRequest
  readonly tools: T
  readonly maxSteps?: number
}

/** Test-owned continuation loop. Production callers must own durable history. */
export const runTools = <T extends Tools>(options: RunOptions<T>) =>
  Stream.unwrap(
    Effect.gen(function* () {
      const names = new Set(Object.keys(options.tools))
      let request = LLMRequest.update(options.request, {
        tools: [...options.request.tools.filter((tool) => !names.has(tool.name)), ...toDefinitions(options.tools)],
      })
      let usage: Usage | undefined
      const events: LLMEvent[] = []

      for (let step = 0; step < (options.maxSteps ?? 10); step++) {
        const streamed = Array.from(yield* LLMClient.stream(request).pipe(Stream.runCollect))
        const state = stepState(streamed)
        usage = addUsage(usage, state.usage)
        events.push(...streamed.filter((event) => event.type !== "finish").map((event) => indexStep(event, step)))

        if (state.toolCalls.length === 0) {
          events.push(LLMEvent.finish({ reason: state.reason, usage, providerMetadata: state.providerMetadata }))
          return Stream.fromIterable(events)
        }

        const dispatched = yield* Effect.forEach(
          state.toolCalls,
          (call) => ToolRuntime.dispatch(options.tools, call).pipe(Effect.map((result) => [call, result] as const)),
          { concurrency: 10 },
        )
        events.push(...dispatched.flatMap(([, result]) => result.events))

        if (step + 1 >= (options.maxSteps ?? 10)) {
          events.push(LLMEvent.finish({ reason: state.reason, usage, providerMetadata: state.providerMetadata }))
          return Stream.fromIterable(events)
        }

        request = LLMRequest.update(request, {
          messages: [
            ...request.messages,
            Message.assistant(state.assistantContent),
            ...dispatched.map(([call, dispatched]) =>
              Message.tool({ id: call.id, name: call.name, result: dispatched.result }),
            ),
          ],
        })
      }

      return Stream.fromIterable(events)
    }),
  )

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
}

const stepState = (events: ReadonlyArray<LLMEvent>) => {
  const assistantContent: ContentPart[] = []
  const toolCalls: ToolCallPart[] = []
  let reason: Extract<LLMEvent, { type: "finish" }>["reason"] = "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)
    } else if (event.type === "tool-call") {
      assistantContent.push(event)
      if (!event.providerExecuted) toolCalls.push(event)
    } else if (event.type === "tool-result" && event.providerExecuted && event.result !== undefined) {
      assistantContent.push(
        ToolResultPart.make({
          id: event.id,
          name: event.name,
          result: event.result,
          providerExecuted: true,
          providerMetadata: event.providerMetadata,
        }),
      )
    } else if (event.type === "finish") {
      reason = event.reason
      usage = event.usage
      providerMetadata = event.providerMetadata
    }
  }
  return { assistantContent, toolCalls, reason, usage, providerMetadata }
}

const appendText = (
  content: ContentPart[],
  type: "text" | "reasoning",
  text: string,
  providerMetadata?: ProviderMetadata,
) => {
  const last = content.at(-1)
  if (last?.type === type) {
    content[content.length - 1] = {
      ...last,
      text: `${last.text}${text}`,
      providerMetadata: providerMetadata ?? last.providerMetadata,
    }
    return
  }
  content.push({ type, text, providerMetadata })
}

const addUsage = (left: Usage | undefined, right: Usage | undefined): Usage | undefined => {
  if (!left) return right
  if (!right) return left
  const sum = (key: keyof Usage) =>
    typeof left[key] !== "number" && typeof right[key] !== "number"
      ? undefined
      : ((left[key] as number | undefined) ?? 0) + ((right[key] as number | undefined) ?? 0)
  return {
    inputTokens: sum("inputTokens"),
    outputTokens: sum("outputTokens"),
    nonCachedInputTokens: sum("nonCachedInputTokens"),
    cacheReadInputTokens: sum("cacheReadInputTokens"),
    cacheWriteInputTokens: sum("cacheWriteInputTokens"),
    reasoningTokens: sum("reasoningTokens"),
    totalTokens: sum("totalTokens"),
  } as Usage
}