@tanstack/ai 0.43.1 → 0.44.0

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 (69) hide show
  1. package/dist/esm/activities/chat/messages.js +21 -8
  2. package/dist/esm/activities/chat/messages.js.map +1 -1
  3. package/dist/esm/activities/embed/adapter.d.ts +69 -0
  4. package/dist/esm/activities/embed/adapter.js +23 -0
  5. package/dist/esm/activities/embed/adapter.js.map +1 -0
  6. package/dist/esm/activities/embed/index.d.ts +117 -0
  7. package/dist/esm/activities/embed/index.js +166 -0
  8. package/dist/esm/activities/embed/index.js.map +1 -0
  9. package/dist/esm/activities/error-payload.d.ts +8 -0
  10. package/dist/esm/activities/error-payload.js +29 -17
  11. package/dist/esm/activities/error-payload.js.map +1 -1
  12. package/dist/esm/activities/generateAudio/index.d.ts +12 -0
  13. package/dist/esm/activities/generateAudio/index.js +19 -6
  14. package/dist/esm/activities/generateAudio/index.js.map +1 -1
  15. package/dist/esm/activities/generateImage/index.d.ts +12 -0
  16. package/dist/esm/activities/generateImage/index.js +21 -7
  17. package/dist/esm/activities/generateImage/index.js.map +1 -1
  18. package/dist/esm/activities/generateSpeech/index.d.ts +17 -1
  19. package/dist/esm/activities/generateSpeech/index.js +19 -6
  20. package/dist/esm/activities/generateSpeech/index.js.map +1 -1
  21. package/dist/esm/activities/generateTranscription/index.d.ts +17 -1
  22. package/dist/esm/activities/generateTranscription/index.js +19 -6
  23. package/dist/esm/activities/generateTranscription/index.js.map +1 -1
  24. package/dist/esm/activities/generateVideo/index.d.ts +18 -0
  25. package/dist/esm/activities/generateVideo/index.js +54 -15
  26. package/dist/esm/activities/generateVideo/index.js.map +1 -1
  27. package/dist/esm/activities/index.d.ts +8 -2
  28. package/dist/esm/activities/index.js +11 -7
  29. package/dist/esm/activities/middleware/types.d.ts +1 -1
  30. package/dist/esm/activities/rerank/adapter.d.ts +63 -0
  31. package/dist/esm/activities/rerank/adapter.js +23 -0
  32. package/dist/esm/activities/rerank/adapter.js.map +1 -0
  33. package/dist/esm/activities/rerank/index.d.ts +92 -0
  34. package/dist/esm/activities/rerank/index.js +163 -0
  35. package/dist/esm/activities/rerank/index.js.map +1 -0
  36. package/dist/esm/activities/summarize/index.d.ts +17 -1
  37. package/dist/esm/activities/summarize/index.js +19 -5
  38. package/dist/esm/activities/summarize/index.js.map +1 -1
  39. package/dist/esm/index.d.ts +7 -2
  40. package/dist/esm/index.js +5 -1
  41. package/dist/esm/middlewares/otel.js +20 -2
  42. package/dist/esm/middlewares/otel.js.map +1 -1
  43. package/dist/esm/types.d.ts +195 -0
  44. package/dist/esm/utilities/activity-abort.d.ts +53 -0
  45. package/dist/esm/utilities/activity-abort.js +150 -0
  46. package/dist/esm/utilities/activity-abort.js.map +1 -0
  47. package/dist/esm/utilities/embedding-input.d.ts +32 -0
  48. package/dist/esm/utilities/embedding-input.js +61 -0
  49. package/dist/esm/utilities/embedding-input.js.map +1 -0
  50. package/package.json +3 -3
  51. package/src/activities/chat/messages.ts +30 -1
  52. package/src/activities/embed/adapter.ts +112 -0
  53. package/src/activities/embed/index.ts +318 -0
  54. package/src/activities/error-payload.ts +41 -9
  55. package/src/activities/generateAudio/index.ts +47 -5
  56. package/src/activities/generateImage/index.ts +48 -5
  57. package/src/activities/generateSpeech/index.ts +52 -9
  58. package/src/activities/generateTranscription/index.ts +52 -9
  59. package/src/activities/generateVideo/index.ts +131 -33
  60. package/src/activities/index.ts +44 -0
  61. package/src/activities/middleware/types.ts +2 -0
  62. package/src/activities/rerank/adapter.ts +90 -0
  63. package/src/activities/rerank/index.ts +302 -0
  64. package/src/activities/summarize/index.ts +59 -19
  65. package/src/index.ts +19 -0
  66. package/src/middlewares/otel.ts +38 -3
  67. package/src/types.ts +219 -0
  68. package/src/utilities/activity-abort.ts +197 -0
  69. package/src/utilities/embedding-input.ts +83 -0
@@ -11,11 +11,18 @@ import { resolveDebugOption } from '../../logger/resolve'
11
11
  import {
12
12
  applyGenerationResultTransforms,
13
13
  createGenerationContext,
14
+ runGenerationAbort,
14
15
  runGenerationError,
15
16
  runGenerationFinish,
16
17
  runGenerationStart,
17
18
  runGenerationUsage,
18
19
  } from '../middleware/run'
20
+ import {
21
+ abortReasonMessage,
22
+ createActivityAbortControls,
23
+ isActivityAbortError,
24
+ raceWithAbort,
25
+ } from '../../utilities/activity-abort'
19
26
  import type { InternalLogger } from '../../logger/internal-logger'
20
27
  import type { DebugOption } from '../../logger/types'
21
28
  import type { GenerationMiddleware } from '../middleware/types'
@@ -36,10 +43,11 @@ export const kind = 'tts' as const
36
43
  /**
37
44
  * Extract provider options from a TTSAdapter via ~types.
38
45
  */
39
- export type TTSProviderOptions<TAdapter> =
40
- TAdapter extends TTSAdapter<any, any>
41
- ? TAdapter['~types']['providerOptions']
42
- : object
46
+ export type TTSProviderOptions<TAdapter> = TAdapter extends {
47
+ '~types': { providerOptions: infer P extends object }
48
+ }
49
+ ? P
50
+ : object
43
51
 
44
52
  // ===========================
45
53
  // Activity Options Type
@@ -92,6 +100,18 @@ export interface TTSActivityOptions<
92
100
  threadId?: string
93
101
  /** Stable run id for correlating this run when persisted. */
94
102
  runId?: string
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
95
115
  }
96
116
 
97
117
  // ===========================
@@ -175,12 +195,18 @@ async function runGenerateSpeech<
175
195
  middleware,
176
196
  threadId,
177
197
  runId,
198
+ timeout,
199
+ abortSignal: callerAbortSignal,
178
200
  ...rest
179
201
  } = options
180
202
  const model = adapter.model
181
203
  const requestId = createId('speech')
182
204
  const startTime = Date.now()
183
205
  const logger: InternalLogger = resolveDebugOption(options.debug)
206
+ const abortControls = createActivityAbortControls({
207
+ timeout,
208
+ abortSignal: callerAbortSignal,
209
+ })
184
210
  const providerName =
185
211
  (adapter as { name?: string; provider?: string }).provider ??
186
212
  (adapter as { name?: string }).name ??
@@ -223,7 +249,16 @@ async function runGenerateSpeech<
223
249
  })
224
250
 
225
251
  try {
226
- const rawResult = await adapter.generateSpeech({ ...rest, model, logger })
252
+ const rawResult = await raceWithAbort(
253
+ adapter.generateSpeech({
254
+ ...rest,
255
+ model,
256
+ logger,
257
+ ...(abortControls.signal ? { abortSignal: abortControls.signal } : {}),
258
+ }),
259
+ abortControls.signal,
260
+ )
261
+ abortControls.clear()
227
262
  const result = await applyGenerationResultTransforms(mwCtx, rawResult)
228
263
  const duration = Date.now() - startTime
229
264
 
@@ -263,6 +298,7 @@ async function runGenerateSpeech<
263
298
 
264
299
  return result
265
300
  } catch (error) {
301
+ abortControls.clear()
266
302
  const duration = Date.now() - startTime
267
303
  const err = error as Error
268
304
  aiEventClient.emit('speech:request:error', {
@@ -274,10 +310,17 @@ async function runGenerateSpeech<
274
310
  modelOptions: rest.modelOptions as Record<string, unknown> | undefined,
275
311
  timestamp: Date.now(),
276
312
  })
277
- await runGenerationError(middleware, mwCtx, {
278
- error,
279
- duration,
280
- })
313
+ if (isActivityAbortError(error, abortControls.signal)) {
314
+ await runGenerationAbort(middleware, mwCtx, {
315
+ reason: abortReasonMessage(error, abortControls.signal),
316
+ duration,
317
+ })
318
+ } else {
319
+ await runGenerationError(middleware, mwCtx, {
320
+ error,
321
+ duration,
322
+ })
323
+ }
281
324
  logger.errors('generateSpeech activity failed', {
282
325
  error,
283
326
  source: 'generateSpeech',
@@ -11,11 +11,18 @@ import { resolveDebugOption } from '../../logger/resolve'
11
11
  import {
12
12
  applyGenerationResultTransforms,
13
13
  createGenerationContext,
14
+ runGenerationAbort,
14
15
  runGenerationError,
15
16
  runGenerationFinish,
16
17
  runGenerationStart,
17
18
  runGenerationUsage,
18
19
  } from '../middleware/run'
20
+ import {
21
+ abortReasonMessage,
22
+ createActivityAbortControls,
23
+ isActivityAbortError,
24
+ raceWithAbort,
25
+ } from '../../utilities/activity-abort'
19
26
  import type { InternalLogger } from '../../logger/internal-logger'
20
27
  import type { DebugOption } from '../../logger/types'
21
28
  import type { GenerationMiddleware } from '../middleware/types'
@@ -40,10 +47,11 @@ export const kind = 'transcription' as const
40
47
  /**
41
48
  * Extract provider options from a TranscriptionAdapter via ~types.
42
49
  */
43
- export type TranscriptionProviderOptions<TAdapter> =
44
- TAdapter extends TranscriptionAdapter<any, any>
45
- ? TAdapter['~types']['providerOptions']
46
- : object
50
+ export type TranscriptionProviderOptions<TAdapter> = TAdapter extends {
51
+ '~types': { providerOptions: infer P extends object }
52
+ }
53
+ ? P
54
+ : object
47
55
 
48
56
  // ===========================
49
57
  // Activity Options Type
@@ -99,6 +107,18 @@ export interface TranscriptionActivityOptions<
99
107
  threadId?: string
100
108
  /** Stable run id for correlating this run when persisted. */
101
109
  runId?: string
110
+ /**
111
+ * Maximum duration of this activity invocation in milliseconds.
112
+ * No SDK-wide default — choose a value suitable for the provider and job.
113
+ * Composed with {@link abortSignal}; the first abort wins.
114
+ */
115
+ timeout?: number
116
+ /**
117
+ * Caller cancellation signal (request disconnects, job/runtime cancellation).
118
+ * Composed with {@link timeout} into an effective signal forwarded to the
119
+ * adapter. Request-specific — not stored on global provider client config.
120
+ */
121
+ abortSignal?: AbortSignal
102
122
  }
103
123
 
104
124
  // ===========================
@@ -211,12 +231,18 @@ async function runGenerateTranscription<
211
231
  middleware,
212
232
  threadId,
213
233
  runId,
234
+ timeout,
235
+ abortSignal: callerAbortSignal,
214
236
  ...rest
215
237
  } = options
216
238
  const model = adapter.model
217
239
  const requestId = createId('transcription')
218
240
  const startTime = Date.now()
219
241
  const logger: InternalLogger = resolveDebugOption(options.debug)
242
+ const abortControls = createActivityAbortControls({
243
+ timeout,
244
+ abortSignal: callerAbortSignal,
245
+ })
220
246
  const providerName =
221
247
  (adapter as { name?: string; provider?: string }).provider ??
222
248
  (adapter as { name?: string }).name ??
@@ -258,7 +284,16 @@ async function runGenerateTranscription<
258
284
  })
259
285
 
260
286
  try {
261
- const rawResult = await adapter.transcribe({ ...rest, model, logger })
287
+ const rawResult = await raceWithAbort(
288
+ adapter.transcribe({
289
+ ...rest,
290
+ model,
291
+ logger,
292
+ ...(abortControls.signal ? { abortSignal: abortControls.signal } : {}),
293
+ }),
294
+ abortControls.signal,
295
+ )
296
+ abortControls.clear()
262
297
  const result = await applyGenerationResultTransforms(mwCtx, rawResult)
263
298
  const duration = Date.now() - startTime
264
299
 
@@ -286,6 +321,7 @@ async function runGenerateTranscription<
286
321
 
287
322
  return result
288
323
  } catch (error) {
324
+ abortControls.clear()
289
325
  const duration = Date.now() - startTime
290
326
  const err = error as Error
291
327
  aiEventClient.emit('transcription:request:error', {
@@ -297,10 +333,17 @@ async function runGenerateTranscription<
297
333
  modelOptions: rest.modelOptions as Record<string, unknown> | undefined,
298
334
  timestamp: Date.now(),
299
335
  })
300
- await runGenerationError(middleware, mwCtx, {
301
- error,
302
- duration,
303
- })
336
+ if (isActivityAbortError(error, abortControls.signal)) {
337
+ await runGenerationAbort(middleware, mwCtx, {
338
+ reason: abortReasonMessage(error, abortControls.signal),
339
+ duration,
340
+ })
341
+ } else {
342
+ await runGenerationError(middleware, mwCtx, {
343
+ error,
344
+ duration,
345
+ })
346
+ }
304
347
  logger.errors('generateTranscription activity failed', {
305
348
  error,
306
349
  source: 'generateTranscription',
@@ -19,6 +19,13 @@ import {
19
19
  runGenerationStart,
20
20
  runGenerationUsage,
21
21
  } from '../middleware/run'
22
+ import {
23
+ abortReasonMessage,
24
+ createActivityAbortControls,
25
+ isActivityAbortError,
26
+ raceWithAbort,
27
+ toAbortError,
28
+ } from '../../utilities/activity-abort'
22
29
  import type { InternalLogger } from '../../logger/internal-logger'
23
30
  import type { DebugOption } from '../../logger/types'
24
31
  import type {
@@ -235,6 +242,24 @@ export type VideoCreateOptions<
235
242
  * job to resume.
236
243
  */
237
244
  middleware?: Array<GenerationMiddleware>
245
+ /**
246
+ * Maximum duration of this activity invocation in milliseconds.
247
+ * No SDK-wide default — choose a value suitable for the provider and job.
248
+ * Composed with {@link abortSignal}; the first abort wins.
249
+ *
250
+ * In stream mode this bounds the full create→poll→complete lifecycle and
251
+ * complements {@link maxDuration} (which defaults to 10 minutes). When both
252
+ * are set, the shorter limit wins via signal composition against the
253
+ * polling deadline.
254
+ */
255
+ timeout?: number
256
+ /**
257
+ * Caller cancellation signal (request disconnects, job/runtime cancellation).
258
+ * Composed with {@link timeout} into an effective signal forwarded to the
259
+ * adapter on job submission. Request-specific — not stored on global
260
+ * provider client config.
261
+ */
262
+ abortSignal?: AbortSignal
238
263
  } & ({} extends VideoProviderOptions<TAdapter>
239
264
  ? {
240
265
  /** Provider-specific options for video generation */ modelOptions?: VideoProviderOptions<TAdapter>
@@ -413,11 +438,24 @@ function videoRunIdForJob(provider: string, jobId: string): string {
413
438
  async function runCreateVideoJob<
414
439
  TAdapter extends VideoAdapter<string, any, any, any, any, any>,
415
440
  >(options: VideoCreateOptions<TAdapter, boolean>): Promise<VideoJobResult> {
416
- const { adapter, prompt, size, duration, modelOptions, middleware } = options
441
+ const {
442
+ adapter,
443
+ prompt,
444
+ size,
445
+ duration,
446
+ modelOptions,
447
+ middleware,
448
+ timeout,
449
+ abortSignal: callerAbortSignal,
450
+ } = options
417
451
  const model = adapter.model
418
452
  const requestId = createId('video')
419
453
  const startTime = Date.now()
420
454
  const logger: InternalLogger = resolveDebugOption(options.debug)
455
+ const abortControls = createActivityAbortControls({
456
+ timeout,
457
+ abortSignal: callerAbortSignal,
458
+ })
421
459
  const providerName =
422
460
  (adapter as { name?: string; provider?: string }).provider ??
423
461
  (adapter as { name?: string }).name ??
@@ -451,24 +489,38 @@ async function runCreateVideoJob<
451
489
 
452
490
  let jobResult: VideoJobResult
453
491
  try {
454
- jobResult = await adapter.createVideoJob({
455
- model,
456
- prompt,
457
- size,
458
- duration,
459
- modelOptions,
460
- logger,
461
- })
492
+ jobResult = await raceWithAbort(
493
+ adapter.createVideoJob({
494
+ model,
495
+ prompt,
496
+ size,
497
+ duration,
498
+ modelOptions,
499
+ logger,
500
+ ...(abortControls.signal ? { abortSignal: abortControls.signal } : {}),
501
+ }),
502
+ abortControls.signal,
503
+ )
504
+ abortControls.clear()
462
505
  } catch (error) {
506
+ abortControls.clear()
463
507
  // No jobId exists, so this run can only be keyed on the request. Start it
464
508
  // just to fail it: `generationRuns.update` on an unknown run id is a no-op
465
509
  // by contract, so without the `onStart` the failure would persist nowhere.
466
510
  const failedCtx = contextFor()
467
511
  await runGenerationStart(middleware, failedCtx)
468
- await runGenerationError(middleware, failedCtx, {
469
- error,
470
- duration: Date.now() - startTime,
471
- })
512
+ const elapsed = Date.now() - startTime
513
+ if (isActivityAbortError(error, abortControls.signal)) {
514
+ await runGenerationAbort(middleware, failedCtx, {
515
+ reason: abortReasonMessage(error, abortControls.signal),
516
+ duration: elapsed,
517
+ })
518
+ } else {
519
+ await runGenerationError(middleware, failedCtx, {
520
+ error,
521
+ duration: elapsed,
522
+ })
523
+ }
472
524
  logger.errors('generateVideo activity failed', {
473
525
  error,
474
526
  source: 'generateVideo',
@@ -489,8 +541,25 @@ async function runCreateVideoJob<
489
541
  return await applyGenerationResultTransforms(mwCtx, jobResult)
490
542
  }
491
543
 
492
- function sleep(ms: number): Promise<void> {
493
- return new Promise((resolve) => setTimeout(resolve, ms))
544
+ function sleep(ms: number, signal?: AbortSignal): Promise<void> {
545
+ if (!signal) {
546
+ return new Promise((resolve) => setTimeout(resolve, ms))
547
+ }
548
+ if (signal.aborted) {
549
+ return Promise.reject(toAbortError(signal.reason))
550
+ }
551
+ return new Promise((resolve, reject) => {
552
+ const timer = setTimeout(() => {
553
+ signal.removeEventListener('abort', onAbort)
554
+ resolve()
555
+ }, ms)
556
+ const onAbort = () => {
557
+ clearTimeout(timer)
558
+ signal.removeEventListener('abort', onAbort)
559
+ reject(toAbortError(signal.reason))
560
+ }
561
+ signal.addEventListener('abort', onAbort, { once: true })
562
+ })
494
563
  }
495
564
 
496
565
  /**
@@ -500,7 +569,16 @@ function sleep(ms: number): Promise<void> {
500
569
  async function* runStreamingVideoGeneration<
501
570
  TAdapter extends VideoAdapter<string, any, any, any, any, any>,
502
571
  >(options: VideoCreateOptions<TAdapter, true>): AsyncIterable<StreamChunk> {
503
- const { adapter, prompt, size, duration, modelOptions, middleware } = options
572
+ const {
573
+ adapter,
574
+ prompt,
575
+ size,
576
+ duration,
577
+ modelOptions,
578
+ middleware,
579
+ timeout,
580
+ abortSignal: callerAbortSignal,
581
+ } = options
504
582
  const model = adapter.model
505
583
  const runId = options.runId ?? createId('run')
506
584
  const requestId = createId('video')
@@ -508,6 +586,10 @@ async function* runStreamingVideoGeneration<
508
586
  const pollingInterval = options.pollingInterval ?? 2000
509
587
  const maxDuration = options.maxDuration ?? 600_000
510
588
  const logger: InternalLogger = resolveDebugOption(options.debug)
589
+ const abortControls = createActivityAbortControls({
590
+ timeout,
591
+ abortSignal: callerAbortSignal,
592
+ })
511
593
  const providerName =
512
594
  (adapter as { name?: string; provider?: string }).provider ??
513
595
  (adapter as { name?: string }).name ??
@@ -555,19 +637,24 @@ async function* runStreamingVideoGeneration<
555
637
  },
556
638
  )
557
639
 
558
- // Tracks whether a terminal observer event (finish/error) has already fired,
559
- // so the `finally` below can fire one on abandonment without double-firing.
640
+ // Tracks whether a terminal observer event (finish/error/abort) has already
641
+ // fired, so the `finally` below can fire one on abandonment without
642
+ // double-firing.
560
643
  let settled = false
561
644
  try {
562
645
  // Create the video generation job
563
- const jobResult = await adapter.createVideoJob({
564
- model,
565
- prompt,
566
- size,
567
- duration,
568
- modelOptions,
569
- logger,
570
- })
646
+ const jobResult = await raceWithAbort(
647
+ adapter.createVideoJob({
648
+ model,
649
+ prompt,
650
+ size,
651
+ duration,
652
+ modelOptions,
653
+ logger,
654
+ ...(abortControls.signal ? { abortSignal: abortControls.signal } : {}),
655
+ }),
656
+ abortControls.signal,
657
+ )
571
658
 
572
659
  yield {
573
660
  type: 'CUSTOM',
@@ -579,7 +666,7 @@ async function* runStreamingVideoGeneration<
579
666
  // Poll for completion
580
667
  const startTime = Date.now()
581
668
  while (Date.now() - startTime < maxDuration) {
582
- await sleep(pollingInterval)
669
+ await sleep(pollingInterval, abortControls.signal)
583
670
 
584
671
  const statusResult = await adapter.getVideoStatus(jobResult.jobId)
585
672
 
@@ -632,6 +719,7 @@ async function* runStreamingVideoGeneration<
632
719
  usage: urlResult.usage,
633
720
  })
634
721
  settled = true
722
+ abortControls.clear()
635
723
 
636
724
  yield {
637
725
  type: 'CUSTOM',
@@ -657,15 +745,24 @@ async function* runStreamingVideoGeneration<
657
745
 
658
746
  throw new Error('Video generation timed out')
659
747
  } catch (error: unknown) {
748
+ abortControls.clear()
660
749
  const payload = toRunErrorPayload(error, 'Video generation failed')
661
- // Mark settled before firing onError: if a user error-hook throws, the
662
- // `finally` below must still not double-fire onAbort over the same op
750
+ // Mark settled before firing terminal hooks: if a user error-hook throws,
751
+ // the `finally` below must still not double-fire onAbort over the same op
663
752
  // (which would mask the original error and end the span twice).
664
753
  settled = true
665
- await runGenerationError(middleware, mwCtx, {
666
- error,
667
- duration: Date.now() - obsStartTime,
668
- })
754
+ const elapsed = Date.now() - obsStartTime
755
+ if (isActivityAbortError(error, abortControls.signal)) {
756
+ await runGenerationAbort(middleware, mwCtx, {
757
+ reason: abortReasonMessage(error, abortControls.signal),
758
+ duration: elapsed,
759
+ })
760
+ } else {
761
+ await runGenerationError(middleware, mwCtx, {
762
+ error,
763
+ duration: elapsed,
764
+ })
765
+ }
669
766
  logger.errors('generateVideo activity failed', {
670
767
  message: payload.message,
671
768
  code: payload.code,
@@ -681,6 +778,7 @@ async function* runStreamingVideoGeneration<
681
778
  timestamp: Date.now(),
682
779
  } as StreamChunk
683
780
  } finally {
781
+ abortControls.clear()
684
782
  if (!settled) {
685
783
  // The consumer abandoned the stream (broke the `for await` loop or
686
784
  // disconnected) before completion, so the generator is being unwound at
@@ -21,6 +21,8 @@ import type { AnyAudioAdapter } from './generateAudio/adapter'
21
21
  import type { AnyVideoAdapter } from './generateVideo/adapter'
22
22
  import type { AnyTTSAdapter } from './generateSpeech/adapter'
23
23
  import type { AnyTranscriptionAdapter } from './generateTranscription/adapter'
24
+ import type { AnyEmbeddingAdapter } from './embed/adapter'
25
+ import type { AnyRerankAdapter } from './rerank/adapter'
24
26
 
25
27
  // ===========================
26
28
  // Chat Activity
@@ -66,6 +68,25 @@ export {
66
68
  type InferTextProviderOptions,
67
69
  } from './summarize/chat-stream-summarize'
68
70
 
71
+ // ===========================
72
+ // Rerank Activity
73
+ // ===========================
74
+
75
+ export {
76
+ kind as rerankKind,
77
+ rerank,
78
+ createRerankOptions,
79
+ type RerankActivityOptions,
80
+ type RerankProviderOptions,
81
+ } from './rerank/index'
82
+
83
+ export {
84
+ BaseRerankAdapter,
85
+ type RerankAdapter,
86
+ type RerankAdapterConfig,
87
+ type AnyRerankAdapter,
88
+ } from './rerank/adapter'
89
+
69
90
  // ===========================
70
91
  // Image Activity
71
92
  // ===========================
@@ -170,6 +191,25 @@ export {
170
191
  type AnyTranscriptionAdapter,
171
192
  } from './generateTranscription/adapter'
172
193
 
194
+ // ===========================
195
+ // Embed Activity
196
+ // ===========================
197
+
198
+ export {
199
+ kind as embeddingKind,
200
+ embed,
201
+ type EmbedOptions,
202
+ type EmbedProviderOptionsForModel,
203
+ type EmbeddingInputForModel,
204
+ } from './embed/index'
205
+
206
+ export {
207
+ BaseEmbeddingAdapter,
208
+ type EmbeddingAdapter,
209
+ type EmbeddingAdapterConfig,
210
+ type AnyEmbeddingAdapter,
211
+ } from './embed/adapter'
212
+
173
213
  // ===========================
174
214
  // Adapter Union Types
175
215
  // ===========================
@@ -183,6 +223,8 @@ export type AIAdapter =
183
223
  | AnyVideoAdapter
184
224
  | AnyTTSAdapter
185
225
  | AnyTranscriptionAdapter
226
+ | AnyEmbeddingAdapter
227
+ | AnyRerankAdapter
186
228
 
187
229
  /** Union type of all adapter kinds */
188
230
  export type AdapterKind =
@@ -193,3 +235,5 @@ export type AdapterKind =
193
235
  | 'video'
194
236
  | 'tts'
195
237
  | 'transcription'
238
+ | 'embedding'
239
+ | 'rerank'
@@ -41,6 +41,8 @@ export type GenerationActivity =
41
41
  | 'audio'
42
42
  | 'tts'
43
43
  | 'transcription'
44
+ | 'embedding'
45
+ | 'rerank'
44
46
  | 'summarize'
45
47
 
46
48
  /**
@@ -0,0 +1,90 @@
1
+ import type { RerankAdapterResult, RerankOptions } from '../../types'
2
+
3
+ /**
4
+ * Configuration for rerank adapter instances
5
+ */
6
+ export interface RerankAdapterConfig {
7
+ apiKey?: string
8
+ baseUrl?: string
9
+ timeout?: number
10
+ headers?: Record<string, string>
11
+ }
12
+
13
+ /**
14
+ * Rerank adapter interface with pre-resolved generics.
15
+ *
16
+ * An adapter is created by a provider function: `provider('model')` → `adapter`
17
+ * All type resolution happens at the provider call site, not in this interface.
18
+ *
19
+ * Generic parameters:
20
+ * - TModel: The specific model name (e.g. 'rerank-v3.5')
21
+ * - TProviderOptions: Provider-specific options (already resolved)
22
+ */
23
+ export interface RerankAdapter<
24
+ TModel extends string = string,
25
+ TProviderOptions extends object = Record<string, unknown>,
26
+ > {
27
+ /** Discriminator for adapter kind */
28
+ readonly kind: 'rerank'
29
+ /** Adapter name identifier */
30
+ readonly name: string
31
+ /** The model this adapter is configured for */
32
+ readonly model: TModel
33
+
34
+ /**
35
+ * @internal Type-only properties for inference. Not assigned at runtime.
36
+ */
37
+ '~types': {
38
+ providerOptions: TProviderOptions
39
+ }
40
+
41
+ /**
42
+ * Rerank the given (pre-serialized) documents against the query, returning
43
+ * scored indices into `options.documents`. The activity layer maps these
44
+ * back to the caller's original documents.
45
+ */
46
+ rerank: (
47
+ options: RerankOptions<TProviderOptions>,
48
+ ) => Promise<RerankAdapterResult>
49
+ }
50
+
51
+ /**
52
+ * A RerankAdapter with any/unknown type parameters.
53
+ * Useful as a constraint in generic functions and interfaces.
54
+ */
55
+ export type AnyRerankAdapter = RerankAdapter<any, any>
56
+
57
+ /**
58
+ * Abstract base class for rerank adapters.
59
+ * Extend this class to implement a rerank adapter for a specific provider.
60
+ *
61
+ * Generic parameters match RerankAdapter - all pre-resolved by the provider function.
62
+ */
63
+ export abstract class BaseRerankAdapter<
64
+ TModel extends string = string,
65
+ TProviderOptions extends object = Record<string, unknown>,
66
+ > implements RerankAdapter<TModel, TProviderOptions> {
67
+ readonly kind = 'rerank' as const
68
+ abstract readonly name: string
69
+ readonly model: TModel
70
+
71
+ // Type-only property - never assigned at runtime
72
+ declare '~types': {
73
+ providerOptions: TProviderOptions
74
+ }
75
+
76
+ protected config: RerankAdapterConfig
77
+
78
+ constructor(config: RerankAdapterConfig = {}, model: TModel) {
79
+ this.config = config
80
+ this.model = model
81
+ }
82
+
83
+ abstract rerank(
84
+ options: RerankOptions<TProviderOptions>,
85
+ ): Promise<RerankAdapterResult>
86
+
87
+ protected generateId(): string {
88
+ return `${this.name}-${Date.now()}-${Math.random().toString(36).slice(2, 9)}`
89
+ }
90
+ }