@tanstack/ai 0.43.1 → 0.44.1

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (73) hide show
  1. package/dist/esm/activities/chat/index.js +18 -0
  2. package/dist/esm/activities/chat/index.js.map +1 -1
  3. package/dist/esm/activities/chat/messages.js +21 -8
  4. package/dist/esm/activities/chat/messages.js.map +1 -1
  5. package/dist/esm/activities/embed/adapter.d.ts +69 -0
  6. package/dist/esm/activities/embed/adapter.js +23 -0
  7. package/dist/esm/activities/embed/adapter.js.map +1 -0
  8. package/dist/esm/activities/embed/index.d.ts +117 -0
  9. package/dist/esm/activities/embed/index.js +166 -0
  10. package/dist/esm/activities/embed/index.js.map +1 -0
  11. package/dist/esm/activities/error-payload.d.ts +8 -0
  12. package/dist/esm/activities/error-payload.js +29 -17
  13. package/dist/esm/activities/error-payload.js.map +1 -1
  14. package/dist/esm/activities/generateAudio/index.d.ts +12 -0
  15. package/dist/esm/activities/generateAudio/index.js +19 -6
  16. package/dist/esm/activities/generateAudio/index.js.map +1 -1
  17. package/dist/esm/activities/generateImage/index.d.ts +12 -0
  18. package/dist/esm/activities/generateImage/index.js +21 -7
  19. package/dist/esm/activities/generateImage/index.js.map +1 -1
  20. package/dist/esm/activities/generateSpeech/index.d.ts +17 -1
  21. package/dist/esm/activities/generateSpeech/index.js +19 -6
  22. package/dist/esm/activities/generateSpeech/index.js.map +1 -1
  23. package/dist/esm/activities/generateTranscription/index.d.ts +17 -1
  24. package/dist/esm/activities/generateTranscription/index.js +19 -6
  25. package/dist/esm/activities/generateTranscription/index.js.map +1 -1
  26. package/dist/esm/activities/generateVideo/index.d.ts +18 -0
  27. package/dist/esm/activities/generateVideo/index.js +54 -15
  28. package/dist/esm/activities/generateVideo/index.js.map +1 -1
  29. package/dist/esm/activities/index.d.ts +8 -2
  30. package/dist/esm/activities/index.js +11 -7
  31. package/dist/esm/activities/middleware/types.d.ts +1 -1
  32. package/dist/esm/activities/rerank/adapter.d.ts +63 -0
  33. package/dist/esm/activities/rerank/adapter.js +23 -0
  34. package/dist/esm/activities/rerank/adapter.js.map +1 -0
  35. package/dist/esm/activities/rerank/index.d.ts +92 -0
  36. package/dist/esm/activities/rerank/index.js +163 -0
  37. package/dist/esm/activities/rerank/index.js.map +1 -0
  38. package/dist/esm/activities/summarize/index.d.ts +17 -1
  39. package/dist/esm/activities/summarize/index.js +19 -5
  40. package/dist/esm/activities/summarize/index.js.map +1 -1
  41. package/dist/esm/index.d.ts +7 -2
  42. package/dist/esm/index.js +5 -1
  43. package/dist/esm/middlewares/otel.js +20 -2
  44. package/dist/esm/middlewares/otel.js.map +1 -1
  45. package/dist/esm/types.d.ts +195 -0
  46. package/dist/esm/utilities/activity-abort.d.ts +53 -0
  47. package/dist/esm/utilities/activity-abort.js +150 -0
  48. package/dist/esm/utilities/activity-abort.js.map +1 -0
  49. package/dist/esm/utilities/embedding-input.d.ts +32 -0
  50. package/dist/esm/utilities/embedding-input.js +61 -0
  51. package/dist/esm/utilities/embedding-input.js.map +1 -0
  52. package/package.json +3 -3
  53. package/skills/ai-core/media-generation/SKILL.md +55 -23
  54. package/src/activities/chat/index.ts +27 -2
  55. package/src/activities/chat/messages.ts +30 -1
  56. package/src/activities/embed/adapter.ts +112 -0
  57. package/src/activities/embed/index.ts +318 -0
  58. package/src/activities/error-payload.ts +41 -9
  59. package/src/activities/generateAudio/index.ts +47 -5
  60. package/src/activities/generateImage/index.ts +48 -5
  61. package/src/activities/generateSpeech/index.ts +52 -9
  62. package/src/activities/generateTranscription/index.ts +52 -9
  63. package/src/activities/generateVideo/index.ts +131 -33
  64. package/src/activities/index.ts +44 -0
  65. package/src/activities/middleware/types.ts +2 -0
  66. package/src/activities/rerank/adapter.ts +90 -0
  67. package/src/activities/rerank/index.ts +302 -0
  68. package/src/activities/summarize/index.ts +59 -19
  69. package/src/index.ts +19 -0
  70. package/src/middlewares/otel.ts +38 -3
  71. package/src/types.ts +219 -0
  72. package/src/utilities/activity-abort.ts +197 -0
  73. package/src/utilities/embedding-input.ts +83 -0
@@ -0,0 +1,302 @@
1
+ /**
2
+ * Rerank Activity
3
+ *
4
+ * Reorders a set of documents by semantic relevance to a query.
5
+ * This is a self-contained module with implementation, types, and JSDoc.
6
+ */
7
+
8
+ import { aiEventClient } from '@tanstack/ai-event-client'
9
+ import { resolveDebugOption } from '../../logger/resolve'
10
+ import { isAbortShapedError } from '../error-payload'
11
+ import {
12
+ createGenerationContext,
13
+ runGenerationAbort,
14
+ runGenerationError,
15
+ runGenerationFinish,
16
+ runGenerationStart,
17
+ runGenerationUsage,
18
+ } from '../middleware/run'
19
+ import type { InternalLogger } from '../../logger/internal-logger'
20
+ import type { DebugOption } from '../../logger/types'
21
+ import type { GenerationMiddleware } from '../middleware/types'
22
+ import type { RerankAdapter } from './adapter'
23
+ import type { RerankResult } from '../../types'
24
+
25
+ // ===========================
26
+ // Activity Kind
27
+ // ===========================
28
+
29
+ /** The adapter kind this activity handles */
30
+ export const kind = 'rerank' as const
31
+
32
+ // ===========================
33
+ // Type Extraction Helpers
34
+ // ===========================
35
+
36
+ /** Extract provider options from a RerankAdapter via ~types */
37
+ export type RerankProviderOptions<TAdapter> = TAdapter extends {
38
+ '~types': { providerOptions: infer P extends object }
39
+ }
40
+ ? P
41
+ : object
42
+
43
+ // ===========================
44
+ // Activity Options Type
45
+ // ===========================
46
+
47
+ /**
48
+ * Options for the rerank activity. The model is extracted from the adapter's
49
+ * model property.
50
+ *
51
+ * @template TAdapter - The rerank adapter type
52
+ * @template TDocument - The document element type (string or object)
53
+ */
54
+ export interface RerankActivityOptions<
55
+ TAdapter extends RerankAdapter<string, RerankProviderOptions<TAdapter>>,
56
+ TDocument extends string | object = string,
57
+ > {
58
+ /** The rerank adapter to use (must be created with a model) */
59
+ adapter: TAdapter & { kind: typeof kind }
60
+ /** The query documents are scored against. */
61
+ query: string
62
+ /**
63
+ * Documents to rerank. Either strings or JSON-serializable objects — object
64
+ * documents are serialized with `JSON.stringify` before being sent to the
65
+ * provider, and the original element (string or object) is returned in the
66
+ * result, preserving its type.
67
+ */
68
+ documents: Array<TDocument>
69
+ /** Return only the top N results. */
70
+ topN?: number
71
+ /** Provider-specific options */
72
+ modelOptions?: RerankProviderOptions<TAdapter>
73
+ /** Forwarded to the provider request for cancellation. */
74
+ abortSignal?: AbortSignal
75
+ /**
76
+ * Observe-only middleware notified on start, usage, success, abort, and
77
+ * error. Pass `otelMiddleware()` to emit OpenTelemetry spans, or implement
78
+ * the `GenerationMiddleware` contract for a custom backend.
79
+ */
80
+ middleware?: Array<GenerationMiddleware>
81
+ /**
82
+ * Enable debug logging. Pass `true` to enable all categories, `false` to
83
+ * silence everything including errors, or a `DebugConfig` object for granular
84
+ * control and/or a custom `Logger`.
85
+ */
86
+ debug?: DebugOption
87
+ }
88
+
89
+ // ===========================
90
+ // Helper Functions
91
+ // ===========================
92
+
93
+ function createId(prefix: string): string {
94
+ return `${prefix}-${Date.now()}-${Math.random().toString(36).slice(2, 9)}`
95
+ }
96
+
97
+ /** Serialize a document for the provider. Strings pass through untouched. */
98
+ function serializeDocument(document: string | object): string {
99
+ return typeof document === 'string' ? document : JSON.stringify(document)
100
+ }
101
+
102
+ function isAbortError(error: unknown, signal?: AbortSignal): boolean {
103
+ // Prefer the error's own identity over the signal state. A genuine
104
+ // cancellation throws an abort-shaped error (DOM `AbortError`, the OpenRouter
105
+ // SDK's `RequestAbortedError`, …). Classifying on `signal.aborted` alone would
106
+ // misroute a real failure — e.g. the out-of-range-index throw below — to the
107
+ // abort hook whenever a shared/long-lived signal happens to already be
108
+ // aborted, hiding it from `onError` observers.
109
+ if (isAbortShapedError(error)) return true
110
+ // Fall back to signal state only for non-Error throws we can't otherwise
111
+ // identify; a real Error with a non-abort name is never an abort.
112
+ return error instanceof Error ? false : signal?.aborted === true
113
+ }
114
+
115
+ // ===========================
116
+ // Activity Implementation
117
+ // ===========================
118
+
119
+ /**
120
+ * Rerank activity - reorders documents by relevance to a query.
121
+ *
122
+ * @example Basic reranking
123
+ * ```ts
124
+ * import { rerank } from '@tanstack/ai'
125
+ * import { cohereRerank } from '@tanstack/ai-cohere'
126
+ *
127
+ * const { ranking, rerankedDocuments } = await rerank({
128
+ * adapter: cohereRerank('rerank-v3.5'),
129
+ * query: 'talk about rain',
130
+ * documents: ['sunny day at the beach', 'rainy afternoon in the city'],
131
+ * topN: 2,
132
+ * })
133
+ *
134
+ * console.log(rerankedDocuments[0]) // 'rainy afternoon in the city'
135
+ * ```
136
+ *
137
+ * @example Reranking object documents
138
+ * ```ts
139
+ * const { ranking } = await rerank({
140
+ * adapter: cohereRerank('rerank-v3.5'),
141
+ * query: 'best laptop for travel',
142
+ * documents: [
143
+ * { id: 1, text: 'A heavy gaming desktop' },
144
+ * { id: 2, text: 'A lightweight ultrabook with all-day battery' },
145
+ * ],
146
+ * })
147
+ *
148
+ * // ranking[0].document is the original object, fully typed.
149
+ * console.log(ranking[0].document.id)
150
+ * ```
151
+ */
152
+ export async function rerank<
153
+ TAdapter extends RerankAdapter<string, RerankProviderOptions<TAdapter>>,
154
+ TDocument extends string | object = string,
155
+ >(
156
+ options: RerankActivityOptions<TAdapter, TDocument>,
157
+ ): Promise<RerankResult<TDocument>> {
158
+ const {
159
+ adapter,
160
+ query,
161
+ documents,
162
+ topN,
163
+ modelOptions,
164
+ abortSignal,
165
+ middleware,
166
+ } = options
167
+ const model = adapter.model
168
+ const requestId = createId('rerank')
169
+ const startTime = Date.now()
170
+ const logger: InternalLogger = resolveDebugOption(options.debug)
171
+
172
+ if (documents.length === 0) {
173
+ throw new Error('rerank() requires at least one document')
174
+ }
175
+
176
+ const mwCtx = createGenerationContext({
177
+ requestId,
178
+ // `rerank` joins the GenerationActivity union; otel maps it to its own
179
+ // gen_ai.operation.name.
180
+ activity: 'rerank',
181
+ provider: adapter.name,
182
+ model,
183
+ modelOptions,
184
+ createId,
185
+ })
186
+
187
+ await runGenerationStart(middleware, mwCtx)
188
+
189
+ aiEventClient.emit('rerank:request:started', {
190
+ requestId,
191
+ provider: adapter.name,
192
+ model,
193
+ documentCount: documents.length,
194
+ timestamp: startTime,
195
+ })
196
+
197
+ logger.request(`activity=rerank provider=${adapter.name}`, {
198
+ provider: adapter.name,
199
+ model,
200
+ documentCount: documents.length,
201
+ })
202
+
203
+ // Serialize once; reuse for the request only. Original documents are mapped
204
+ // back by index below so the caller's element type is preserved.
205
+ const serialized = documents.map(serializeDocument)
206
+
207
+ try {
208
+ const result = await adapter.rerank({
209
+ model,
210
+ query,
211
+ documents: serialized,
212
+ topN,
213
+ modelOptions,
214
+ abortSignal,
215
+ logger,
216
+ })
217
+
218
+ const ranking = result.ranking.map((r) => {
219
+ const document = documents[r.index]
220
+ if (document === undefined) {
221
+ throw new Error(
222
+ `rerank(): provider ${adapter.name} returned out-of-range index ${r.index}`,
223
+ )
224
+ }
225
+ return { index: r.index, score: r.score, document }
226
+ })
227
+ const rerankedDocuments = ranking.map((r) => r.document)
228
+
229
+ const duration = Date.now() - startTime
230
+
231
+ aiEventClient.emit('rerank:request:completed', {
232
+ requestId,
233
+ provider: adapter.name,
234
+ model,
235
+ documentCount: documents.length,
236
+ resultCount: ranking.length,
237
+ duration,
238
+ timestamp: Date.now(),
239
+ })
240
+
241
+ aiEventClient.emit('rerank:usage', {
242
+ requestId,
243
+ model,
244
+ usage: result.usage,
245
+ timestamp: Date.now(),
246
+ })
247
+
248
+ logger.output(`activity=rerank results=${ranking.length}`, {
249
+ resultCount: ranking.length,
250
+ })
251
+
252
+ await runGenerationUsage(middleware, mwCtx, result.usage)
253
+ await runGenerationFinish(middleware, mwCtx, {
254
+ duration,
255
+ usage: result.usage,
256
+ })
257
+
258
+ return {
259
+ id: result.id,
260
+ model,
261
+ ranking,
262
+ rerankedDocuments,
263
+ usage: result.usage,
264
+ }
265
+ } catch (error) {
266
+ const duration = Date.now() - startTime
267
+ if (isAbortError(error, abortSignal)) {
268
+ await runGenerationAbort(middleware, mwCtx, {
269
+ reason: error instanceof Error ? error.message : undefined,
270
+ duration,
271
+ })
272
+ } else {
273
+ await runGenerationError(middleware, mwCtx, { error, duration })
274
+ }
275
+ logger.errors('rerank activity failed', { error, source: 'rerank' })
276
+ throw error
277
+ }
278
+ }
279
+
280
+ // ===========================
281
+ // Options Factory
282
+ // ===========================
283
+
284
+ /**
285
+ * Create typed options for the rerank() function without executing.
286
+ */
287
+ export function createRerankOptions<
288
+ TAdapter extends RerankAdapter<string, RerankProviderOptions<TAdapter>>,
289
+ TDocument extends string | object = string,
290
+ >(
291
+ options: RerankActivityOptions<TAdapter, TDocument>,
292
+ ): RerankActivityOptions<TAdapter, TDocument> {
293
+ return options
294
+ }
295
+
296
+ // Re-export adapter types
297
+ export type {
298
+ RerankAdapter,
299
+ RerankAdapterConfig,
300
+ AnyRerankAdapter,
301
+ } from './adapter'
302
+ export { BaseRerankAdapter } from './adapter'
@@ -17,6 +17,12 @@ import {
17
17
  runGenerationStart,
18
18
  runGenerationUsage,
19
19
  } from '../middleware/run'
20
+ import {
21
+ abortReasonMessage,
22
+ createActivityAbortControls,
23
+ isActivityAbortError,
24
+ raceWithAbort,
25
+ } from '../../utilities/activity-abort'
20
26
  import type { InternalLogger } from '../../logger/internal-logger'
21
27
  import type { DebugOption } from '../../logger/types'
22
28
  import type { GenerationMiddleware } from '../middleware/types'
@@ -35,10 +41,11 @@ export const kind = 'summarize' as const
35
41
  // ===========================
36
42
 
37
43
  /** Extract provider options from a SummarizeAdapter via ~types */
38
- export type SummarizeProviderOptions<TAdapter> =
39
- TAdapter extends SummarizeAdapter<any, any>
40
- ? TAdapter['~types']['providerOptions']
41
- : object
44
+ export type SummarizeProviderOptions<TAdapter> = TAdapter extends {
45
+ '~types': { providerOptions: infer P extends object }
46
+ }
47
+ ? P
48
+ : object
42
49
 
43
50
  // ===========================
44
51
  // Activity Options Type
@@ -93,6 +100,18 @@ export interface SummarizeActivityOptions<
93
100
  * mid-summary fires `onAbort`.
94
101
  */
95
102
  middleware?: Array<GenerationMiddleware>
103
+ /**
104
+ * Maximum duration of this activity invocation in milliseconds.
105
+ * No SDK-wide default — choose a value suitable for the provider and job.
106
+ * Composed with {@link abortSignal}; the first abort wins.
107
+ */
108
+ timeout?: number
109
+ /**
110
+ * Caller cancellation signal (request disconnects, job/runtime cancellation).
111
+ * Composed with {@link timeout} into an effective signal forwarded to the
112
+ * adapter. Request-specific — not stored on global provider client config.
113
+ */
114
+ abortSignal?: AbortSignal
96
115
  /**
97
116
  * Whether to stream the summarization result.
98
117
  * When true, returns an AsyncIterable<StreamChunk> for streaming output.
@@ -195,18 +214,12 @@ export function summarize<
195
214
 
196
215
  if (stream) {
197
216
  return runStreamingSummarize(
198
- options as SummarizeActivityOptions<
199
- SummarizeAdapter<string, object>,
200
- true
201
- >,
217
+ options as SummarizeActivityOptions<TAdapter, true>,
202
218
  ) as SummarizeActivityResult<TStream>
203
219
  }
204
220
 
205
221
  return runSummarize(
206
- options as SummarizeActivityOptions<
207
- SummarizeAdapter<string, object>,
208
- false
209
- >,
222
+ options as SummarizeActivityOptions<TAdapter, false>,
210
223
  ) as SummarizeActivityResult<TStream>
211
224
  }
212
225
 
@@ -216,13 +229,26 @@ export function summarize<
216
229
  async function runSummarize(
217
230
  options: SummarizeActivityOptions<SummarizeAdapter<string, object>, false>,
218
231
  ): Promise<SummarizationResult> {
219
- const { adapter, text, maxLength, style, focus, modelOptions, middleware } =
220
- options
232
+ const {
233
+ adapter,
234
+ text,
235
+ maxLength,
236
+ style,
237
+ focus,
238
+ modelOptions,
239
+ middleware,
240
+ timeout,
241
+ abortSignal: callerAbortSignal,
242
+ } = options
221
243
  const model = adapter.model
222
244
  const requestId = createId('summarize')
223
245
  const inputLength = text.length
224
246
  const startTime = Date.now()
225
247
  const logger: InternalLogger = resolveDebugOption(options.debug)
248
+ const abortControls = createActivityAbortControls({
249
+ timeout,
250
+ abortSignal: callerAbortSignal,
251
+ })
226
252
 
227
253
  const mwCtx = createGenerationContext({
228
254
  requestId,
@@ -259,10 +285,15 @@ async function runSummarize(
259
285
  focus,
260
286
  modelOptions,
261
287
  logger,
288
+ ...(abortControls.signal ? { abortSignal: abortControls.signal } : {}),
262
289
  }
263
290
 
264
291
  try {
265
- const rawResult = await adapter.summarize(summarizeOptions)
292
+ const rawResult = await raceWithAbort(
293
+ adapter.summarize(summarizeOptions),
294
+ abortControls.signal,
295
+ )
296
+ abortControls.clear()
266
297
  // Transforms run before anything observes the result — the same order every
267
298
  // media activity uses — so the run record and the returned value are the
268
299
  // same object.
@@ -294,10 +325,19 @@ async function runSummarize(
294
325
 
295
326
  return result
296
327
  } catch (error) {
297
- await runGenerationError(middleware, mwCtx, {
298
- error,
299
- duration: Date.now() - startTime,
300
- })
328
+ abortControls.clear()
329
+ const duration = Date.now() - startTime
330
+ if (isActivityAbortError(error, abortControls.signal)) {
331
+ await runGenerationAbort(middleware, mwCtx, {
332
+ reason: abortReasonMessage(error, abortControls.signal),
333
+ duration,
334
+ })
335
+ } else {
336
+ await runGenerationError(middleware, mwCtx, {
337
+ error,
338
+ duration,
339
+ })
340
+ }
301
341
  logger.errors('summarize activity failed', {
302
342
  error,
303
343
  source: 'summarize',
package/src/index.ts CHANGED
@@ -2,22 +2,26 @@
2
2
  export {
3
3
  chat,
4
4
  summarize,
5
+ rerank,
5
6
  generateImage,
6
7
  generateAudio,
7
8
  generateVideo,
8
9
  getVideoJobStatus,
9
10
  generateSpeech,
10
11
  generateTranscription,
12
+ embed,
11
13
  } from './activities/index'
12
14
 
13
15
  // Create options functions - for pre-defining typed configurations
14
16
  export { createChatOptions } from './activities/chat/index'
15
17
  export { createSummarizeOptions } from './activities/summarize/index'
18
+ export { createRerankOptions } from './activities/rerank/index'
16
19
  export { createImageOptions } from './activities/generateImage/index'
17
20
  export { createAudioOptions } from './activities/generateAudio/index'
18
21
  export { createVideoOptions } from './activities/generateVideo/index'
19
22
  export { createSpeechOptions } from './activities/generateSpeech/index'
20
23
  export { createTranscriptionOptions } from './activities/generateTranscription/index'
24
+ export { createEmbedOptions } from './activities/embed/index'
21
25
 
22
26
  // Re-export types
23
27
  export type {
@@ -36,8 +40,15 @@ export type {
36
40
  TranscriptionAdapter,
37
41
  AnyVideoAdapter,
38
42
  VideoAdapter,
43
+ AnyEmbeddingAdapter,
44
+ EmbeddingAdapter,
45
+ AnyRerankAdapter,
46
+ RerankAdapter,
39
47
  } from './activities/index'
40
48
 
49
+ // Rerank adapter base + types
50
+ export { BaseRerankAdapter } from './activities/rerank/adapter'
51
+
41
52
  // Tool definition
42
53
  export {
43
54
  toolDefinition,
@@ -303,6 +314,14 @@ export { buildBaseUsage, type BaseUsageInput } from './utilities/usage'
303
314
  export { resolveMediaPrompt } from './utilities/media-prompt'
304
315
  export type { ResolvedMediaPrompt } from './utilities/media-prompt'
305
316
 
317
+ // Embedding input resolution (used by embedding adapters)
318
+ export {
319
+ resolveEmbeddingInput,
320
+ requireTextOnlyEmbeddingInput,
321
+ countEmbeddingInputModalities,
322
+ } from './utilities/embedding-input'
323
+ export type { ResolvedEmbeddingItem } from './utilities/embedding-input'
324
+
306
325
  // System prompts (type + normaliser used by adapters)
307
326
  export type { SystemPrompt, NormalizedSystemPrompt } from './system-prompts'
308
327
  export { normalizeSystemPrompts } from './system-prompts'
@@ -28,6 +28,7 @@ import type {
28
28
  GenerationMiddleware,
29
29
  GenerationMiddlewareContext,
30
30
  } from '../activities/middleware/types'
31
+ import type { TokenUsage } from '../types'
31
32
 
32
33
  /**
33
34
  * Scope (role) of an OTel span emitted by this middleware.
@@ -85,6 +86,8 @@ const OPERATION_NAME: Record<GenerationActivity, string> = {
85
86
  audio: 'audio_generation',
86
87
  tts: 'text_to_speech',
87
88
  transcription: 'transcription',
89
+ embedding: 'embeddings',
90
+ rerank: 'rerank',
88
91
  summarize: 'summarize',
89
92
  }
90
93
 
@@ -141,6 +144,8 @@ interface RequestState {
141
144
  * from the (base-shaped) finish info, which doesn't carry it.
142
145
  */
143
146
  lastFinishReason: string | null
147
+ rootUsageAttributes: Record<string, number> | null
148
+ rootUsageApplied: boolean
144
149
  }
145
150
 
146
151
  const stateByCtx = new WeakMap<ChatMiddlewareContext, RequestState>()
@@ -148,6 +153,29 @@ const stateByCtx = new WeakMap<ChatMiddlewareContext, RequestState>()
148
153
  const DEFAULT_MAX_CONTENT_LENGTH = 100_000
149
154
  const REDACTION_FAILED_SENTINEL = '[redaction_failed]'
150
155
 
156
+ function accumulateUsageAttributes(
157
+ current: Record<string, number> | null,
158
+ usage: TokenUsage,
159
+ ): Record<string, number> {
160
+ const accumulated = current ?? {}
161
+ for (const [key, value] of Object.entries(usageAttributes(usage))) {
162
+ if (typeof value === 'number') {
163
+ accumulated[key] = (accumulated[key] ?? 0) + value
164
+ }
165
+ }
166
+ return accumulated
167
+ }
168
+
169
+ function applyRootUsage(state: RequestState, fallbackUsage?: TokenUsage): void {
170
+ if (state.rootUsageApplied) return
171
+
172
+ const attributes =
173
+ state.rootUsageAttributes ??
174
+ (fallbackUsage ? usageAttributes(fallbackUsage) : null)
175
+ if (attributes) state.rootSpan.setAttributes(attributes)
176
+ state.rootUsageApplied = true
177
+ }
178
+
151
179
  function serializeContent(content: unknown): string {
152
180
  if (typeof content === 'string') return content
153
181
  if (!Array.isArray(content)) return ''
@@ -398,6 +426,8 @@ export function otelMiddleware(
398
426
  assistantTextBufferTruncated: false,
399
427
  startTime: Date.now(),
400
428
  lastFinishReason: null,
429
+ rootUsageAttributes: null,
430
+ rootUsageApplied: false,
401
431
  })
402
432
  })
403
433
  },
@@ -668,6 +698,11 @@ export function otelMiddleware(
668
698
  const state = stateByCtx.get(chatCtx)
669
699
  if (!state) return
670
700
 
701
+ state.rootUsageAttributes = accumulateUsageAttributes(
702
+ state.rootUsageAttributes,
703
+ usage,
704
+ )
705
+
671
706
  // Always record the token histogram — metrics don't depend on having
672
707
  // an iteration span, and skipping here would drop metric data if an
673
708
  // adapter emits `onUsage` outside the iteration window.
@@ -907,6 +942,7 @@ export function otelMiddleware(
907
942
  })
908
943
  }
909
944
 
945
+ applyRootUsage(state)
910
946
  safeCall('otel.onSpanEnd', () =>
911
947
  onSpanEnd?.({ kind: 'chat', ctx: chatCtx }, state.rootSpan),
912
948
  )
@@ -988,6 +1024,7 @@ export function otelMiddleware(
988
1024
  })
989
1025
  }
990
1026
 
1027
+ applyRootUsage(state)
991
1028
  safeCall('otel.onSpanEnd', () =>
992
1029
  onSpanEnd?.({ kind: 'chat', ctx: chatCtx }, state.rootSpan),
993
1030
  )
@@ -1045,9 +1082,6 @@ export function otelMiddleware(
1045
1082
  })
1046
1083
  }
1047
1084
 
1048
- if (info.usage) {
1049
- state.rootSpan.setAttributes(usageAttributes(info.usage))
1050
- }
1051
1085
  if (state.lastFinishReason) {
1052
1086
  state.rootSpan.setAttribute('gen_ai.response.finish_reasons', [
1053
1087
  state.lastFinishReason,
@@ -1058,6 +1092,7 @@ export function otelMiddleware(
1058
1092
  state.iterationCount,
1059
1093
  )
1060
1094
 
1095
+ applyRootUsage(state, info.usage)
1061
1096
  safeCall('otel.onSpanEnd', () =>
1062
1097
  onSpanEnd?.({ kind: 'chat', ctx: chatCtx }, state.rootSpan),
1063
1098
  )