@tanstack/ai 0.6.3 → 0.8.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 (59) hide show
  1. package/dist/esm/activities/chat/index.d.ts +20 -0
  2. package/dist/esm/activities/chat/index.js +248 -213
  3. package/dist/esm/activities/chat/index.js.map +1 -1
  4. package/dist/esm/activities/chat/middleware/compose.d.ts +66 -0
  5. package/dist/esm/activities/chat/middleware/compose.js +327 -0
  6. package/dist/esm/activities/chat/middleware/compose.js.map +1 -0
  7. package/dist/esm/activities/chat/middleware/index.d.ts +2 -0
  8. package/dist/esm/activities/chat/middleware/tool-cache-middleware.d.ts +89 -0
  9. package/dist/esm/activities/chat/middleware/tool-cache-middleware.js +76 -0
  10. package/dist/esm/activities/chat/middleware/tool-cache-middleware.js.map +1 -0
  11. package/dist/esm/activities/chat/middleware/types.d.ts +307 -0
  12. package/dist/esm/activities/chat/tools/tool-calls.d.ts +16 -1
  13. package/dist/esm/activities/chat/tools/tool-calls.js +148 -64
  14. package/dist/esm/activities/chat/tools/tool-calls.js.map +1 -1
  15. package/dist/esm/activities/generateImage/index.js +1 -1
  16. package/dist/esm/activities/generateImage/index.js.map +1 -1
  17. package/dist/esm/activities/generateSpeech/index.js +1 -1
  18. package/dist/esm/activities/generateSpeech/index.js.map +1 -1
  19. package/dist/esm/activities/generateTranscription/index.js +1 -1
  20. package/dist/esm/activities/generateTranscription/index.js.map +1 -1
  21. package/dist/esm/activities/generateVideo/index.js +1 -1
  22. package/dist/esm/activities/generateVideo/index.js.map +1 -1
  23. package/dist/esm/activities/summarize/index.js +1 -1
  24. package/dist/esm/activities/summarize/index.js.map +1 -1
  25. package/dist/esm/index.d.ts +3 -1
  26. package/dist/esm/index.js +2 -2
  27. package/dist/esm/middlewares/content-guard.d.ts +77 -0
  28. package/dist/esm/middlewares/content-guard.js +155 -0
  29. package/dist/esm/middlewares/content-guard.js.map +1 -0
  30. package/dist/esm/middlewares/index.d.ts +2 -0
  31. package/dist/esm/middlewares/index.js +7 -0
  32. package/dist/esm/middlewares/index.js.map +1 -0
  33. package/dist/esm/middlewares/tool-cache.d.ts +1 -0
  34. package/dist/esm/realtime/index.d.ts +30 -0
  35. package/dist/esm/realtime/index.js +8 -0
  36. package/dist/esm/realtime/index.js.map +1 -0
  37. package/dist/esm/realtime/types.d.ts +234 -0
  38. package/package.json +6 -6
  39. package/src/activities/chat/index.ts +322 -256
  40. package/src/activities/chat/middleware/compose.ts +392 -0
  41. package/src/activities/chat/middleware/index.ts +17 -0
  42. package/src/activities/chat/middleware/tool-cache-middleware.ts +189 -0
  43. package/src/activities/chat/middleware/types.ts +419 -0
  44. package/src/activities/chat/tools/tool-calls.ts +225 -87
  45. package/src/activities/generateImage/index.ts +1 -1
  46. package/src/activities/generateSpeech/index.ts +1 -1
  47. package/src/activities/generateTranscription/index.ts +1 -1
  48. package/src/activities/generateVideo/index.ts +1 -1
  49. package/src/activities/summarize/index.ts +1 -1
  50. package/src/index.ts +41 -2
  51. package/src/middlewares/content-guard.ts +285 -0
  52. package/src/middlewares/index.ts +13 -0
  53. package/src/middlewares/tool-cache.ts +6 -0
  54. package/src/realtime/index.ts +38 -0
  55. package/src/realtime/types.ts +294 -0
  56. package/dist/esm/event-client.d.ts +0 -394
  57. package/dist/esm/event-client.js +0 -13
  58. package/dist/esm/event-client.js.map +0 -1
  59. package/src/event-client.ts +0 -497
@@ -10,6 +10,41 @@ import type {
10
10
  ToolCallStartEvent,
11
11
  ToolExecutionContext,
12
12
  } from '../../../types'
13
+ import type {
14
+ AfterToolCallInfo,
15
+ BeforeToolCallDecision,
16
+ } from '../middleware/types'
17
+
18
+ function safeJsonParse(value: string): unknown {
19
+ try {
20
+ return JSON.parse(value)
21
+ } catch {
22
+ return value
23
+ }
24
+ }
25
+
26
+ /**
27
+ * Optional middleware hooks for tool execution.
28
+ * When provided, these callbacks are invoked before/after each tool execution.
29
+ */
30
+ export interface ToolExecutionMiddlewareHooks {
31
+ onBeforeToolCall?: (
32
+ toolCall: ToolCall,
33
+ tool: Tool | undefined,
34
+ args: unknown,
35
+ ) => Promise<BeforeToolCallDecision>
36
+ onAfterToolCall?: (info: AfterToolCallInfo) => Promise<void>
37
+ }
38
+
39
+ /**
40
+ * Error thrown when middleware decides to abort the chat run during tool execution.
41
+ */
42
+ export class MiddlewareAbortError extends Error {
43
+ constructor(reason: string) {
44
+ super(reason)
45
+ this.name = 'MiddlewareAbortError'
46
+ }
47
+ }
13
48
 
14
49
  /**
15
50
  * Manages tool call accumulation and execution for the chat() method's automatic tool execution loop.
@@ -289,6 +324,152 @@ async function* executeWithEventPolling<T>(
289
324
  return state.result
290
325
  }
291
326
 
327
+ /**
328
+ * Apply a middleware onBeforeToolCall decision.
329
+ * Returns the (possibly transformed) input if execution should proceed,
330
+ * or undefined if the tool call was skipped (result already pushed).
331
+ * Throws MiddlewareAbortError if the decision is 'abort'.
332
+ */
333
+ async function applyBeforeToolCallDecision(
334
+ toolCall: ToolCall,
335
+ tool: Tool,
336
+ input: unknown,
337
+ toolName: string,
338
+ middlewareHooks: ToolExecutionMiddlewareHooks,
339
+ results: Array<ToolResult>,
340
+ ): Promise<{ proceed: true; input: unknown } | { proceed: false }> {
341
+ if (!middlewareHooks.onBeforeToolCall) {
342
+ return { proceed: true, input }
343
+ }
344
+
345
+ const decision = await middlewareHooks.onBeforeToolCall(toolCall, tool, input)
346
+ if (!decision) {
347
+ return { proceed: true, input }
348
+ }
349
+
350
+ if (decision.type === 'abort') {
351
+ throw new MiddlewareAbortError(decision.reason || 'Aborted by middleware')
352
+ }
353
+
354
+ if (decision.type === 'skip') {
355
+ const skipResult = decision.result
356
+ results.push({
357
+ toolCallId: toolCall.id,
358
+ toolName,
359
+ result:
360
+ typeof skipResult === 'string'
361
+ ? safeJsonParse(skipResult)
362
+ : skipResult || null,
363
+ duration: 0,
364
+ })
365
+ if (middlewareHooks.onAfterToolCall) {
366
+ await middlewareHooks.onAfterToolCall({
367
+ toolCall,
368
+ tool,
369
+ toolName,
370
+ toolCallId: toolCall.id,
371
+ ok: true,
372
+ duration: 0,
373
+ result: skipResult,
374
+ })
375
+ }
376
+ return { proceed: false }
377
+ }
378
+
379
+ return { proceed: true, input: decision.args }
380
+ }
381
+
382
+ /**
383
+ * Execute a server-side tool with event polling, output validation, and middleware hooks.
384
+ * Yields CustomEvent chunks during execution and pushes the result to the results array.
385
+ */
386
+ async function* executeServerTool(
387
+ toolCall: ToolCall,
388
+ tool: Tool,
389
+ toolName: string,
390
+ input: unknown,
391
+ context: ToolExecutionContext,
392
+ pendingEvents: Array<CustomEvent>,
393
+ results: Array<ToolResult>,
394
+ middlewareHooks?: ToolExecutionMiddlewareHooks,
395
+ ): AsyncGenerator<CustomEvent, void, void> {
396
+ const startTime = Date.now()
397
+ try {
398
+ const executionPromise = Promise.resolve(tool.execute!(input, context))
399
+ let result = yield* executeWithEventPolling(executionPromise, pendingEvents)
400
+ const duration = Date.now() - startTime
401
+
402
+ // Flush remaining events
403
+ while (pendingEvents.length > 0) {
404
+ yield pendingEvents.shift()!
405
+ }
406
+
407
+ // Validate output against outputSchema if provided
408
+ if (
409
+ tool.outputSchema &&
410
+ isStandardSchema(tool.outputSchema) &&
411
+ result !== undefined &&
412
+ result !== null
413
+ ) {
414
+ result = parseWithStandardSchema(tool.outputSchema, result)
415
+ }
416
+
417
+ const finalResult =
418
+ typeof result === 'string' ? safeJsonParse(result) : result || null
419
+
420
+ results.push({
421
+ toolCallId: toolCall.id,
422
+ toolName,
423
+ result: finalResult,
424
+ duration,
425
+ })
426
+
427
+ if (middlewareHooks?.onAfterToolCall) {
428
+ await middlewareHooks.onAfterToolCall({
429
+ toolCall,
430
+ tool,
431
+ toolName,
432
+ toolCallId: toolCall.id,
433
+ ok: true,
434
+ duration,
435
+ result: finalResult,
436
+ })
437
+ }
438
+ } catch (error: unknown) {
439
+ const duration = Date.now() - startTime
440
+
441
+ // Flush remaining events
442
+ while (pendingEvents.length > 0) {
443
+ yield pendingEvents.shift()!
444
+ }
445
+
446
+ if (error instanceof MiddlewareAbortError) {
447
+ throw error
448
+ }
449
+
450
+ const message = error instanceof Error ? error.message : 'Unknown error'
451
+ results.push({
452
+ toolCallId: toolCall.id,
453
+ toolName,
454
+ result: { error: message },
455
+ state: 'output-error',
456
+ duration,
457
+ })
458
+
459
+ if (middlewareHooks?.onAfterToolCall) {
460
+ await middlewareHooks.onAfterToolCall({
461
+ toolCall,
462
+ tool,
463
+ toolName,
464
+ toolCallId: toolCall.id,
465
+ ok: false,
466
+ duration,
467
+ error,
468
+ })
469
+ }
470
+ }
471
+ }
472
+
292
473
  /**
293
474
  * Execute tool calls based on their configuration.
294
475
  * Yields CustomEvent chunks during tool execution for real-time progress updates.
@@ -313,6 +494,7 @@ export async function* executeToolCalls(
313
494
  eventName: string,
314
495
  value: Record<string, any>,
315
496
  ) => CustomEvent,
497
+ middlewareHooks?: ToolExecutionMiddlewareHooks,
316
498
  ): AsyncGenerator<CustomEvent, ExecuteToolCallsResult, void> {
317
499
  const results: Array<ToolResult> = []
318
500
  const needsApproval: Array<ApprovalRequest> = []
@@ -402,13 +584,6 @@ export async function* executeToolCalls(
402
584
  },
403
585
  }
404
586
 
405
- // Helper to flush any pending events
406
- function* flushEvents(): Generator<CustomEvent> {
407
- while (pendingEvents.length > 0) {
408
- yield pendingEvents.shift()!
409
- }
410
- }
411
-
412
587
  // CASE 1: Client-side tool (no execute function)
413
588
  if (!tool.execute) {
414
589
  // Check if tool needs approval
@@ -482,51 +657,30 @@ export async function* executeToolCalls(
482
657
  const approved = approvals.get(approvalId)
483
658
 
484
659
  if (approved) {
485
- // Execute after approval
486
- const startTime = Date.now()
487
- try {
488
- const executionPromise = Promise.resolve(
489
- tool.execute(input, context),
490
- )
491
- let result = yield* executeWithEventPolling(
492
- executionPromise,
493
- pendingEvents,
494
- )
495
- const duration = Date.now() - startTime
496
- yield* flushEvents()
497
-
498
- // Validate output against outputSchema if provided (for Standard Schema compliant schemas)
499
- if (
500
- tool.outputSchema &&
501
- isStandardSchema(tool.outputSchema) &&
502
- result !== undefined &&
503
- result !== null
504
- ) {
505
- result = parseWithStandardSchema(tool.outputSchema, result)
506
- }
507
-
508
- results.push({
509
- toolCallId: toolCall.id,
510
- toolName,
511
- result:
512
- typeof result === 'string'
513
- ? JSON.parse(result)
514
- : result || null,
515
- duration,
516
- })
517
- } catch (error: unknown) {
518
- const duration = Date.now() - startTime
519
- yield* flushEvents()
520
- const message =
521
- error instanceof Error ? error.message : 'Unknown error'
522
- results.push({
523
- toolCallId: toolCall.id,
660
+ // Apply middleware before-hook for approved tools
661
+ if (middlewareHooks) {
662
+ const decision = await applyBeforeToolCallDecision(
663
+ toolCall,
664
+ tool,
665
+ input,
524
666
  toolName,
525
- result: { error: message },
526
- state: 'output-error',
527
- duration,
528
- })
667
+ middlewareHooks,
668
+ results,
669
+ )
670
+ if (!decision.proceed) continue
671
+ input = decision.input
529
672
  }
673
+
674
+ yield* executeServerTool(
675
+ toolCall,
676
+ tool,
677
+ toolName,
678
+ input,
679
+ context,
680
+ pendingEvents,
681
+ results,
682
+ middlewareHooks,
683
+ )
530
684
  } else {
531
685
  // User declined
532
686
  results.push({
@@ -549,45 +703,29 @@ export async function* executeToolCalls(
549
703
  }
550
704
 
551
705
  // CASE 3: Normal server tool - execute immediately
552
- const startTime = Date.now()
553
- try {
554
- const executionPromise = Promise.resolve(tool.execute(input, context))
555
- let result = yield* executeWithEventPolling(
556
- executionPromise,
557
- pendingEvents,
558
- )
559
- const duration = Date.now() - startTime
560
- yield* flushEvents()
561
-
562
- // Validate output against outputSchema if provided (for Standard Schema compliant schemas)
563
- if (
564
- tool.outputSchema &&
565
- isStandardSchema(tool.outputSchema) &&
566
- result !== undefined &&
567
- result !== null
568
- ) {
569
- result = parseWithStandardSchema(tool.outputSchema, result)
570
- }
571
-
572
- results.push({
573
- toolCallId: toolCall.id,
574
- toolName,
575
- result:
576
- typeof result === 'string' ? JSON.parse(result) : result || null,
577
- duration,
578
- })
579
- } catch (error: unknown) {
580
- const duration = Date.now() - startTime
581
- yield* flushEvents()
582
- const message = error instanceof Error ? error.message : 'Unknown error'
583
- results.push({
584
- toolCallId: toolCall.id,
706
+ if (middlewareHooks) {
707
+ const decision = await applyBeforeToolCallDecision(
708
+ toolCall,
709
+ tool,
710
+ input,
585
711
  toolName,
586
- result: { error: message },
587
- state: 'output-error',
588
- duration,
589
- })
712
+ middlewareHooks,
713
+ results,
714
+ )
715
+ if (!decision.proceed) continue
716
+ input = decision.input
590
717
  }
718
+
719
+ yield* executeServerTool(
720
+ toolCall,
721
+ tool,
722
+ toolName,
723
+ input,
724
+ context,
725
+ pendingEvents,
726
+ results,
727
+ middlewareHooks,
728
+ )
591
729
  }
592
730
 
593
731
  return { results, needsApproval, needsClientExecution }
@@ -5,7 +5,7 @@
5
5
  * This is a self-contained module with implementation, types, and JSDoc.
6
6
  */
7
7
 
8
- import { aiEventClient } from '../../event-client.js'
8
+ import { aiEventClient } from '@tanstack/ai-event-client'
9
9
  import { streamGenerationResult } from '../stream-generation-result.js'
10
10
  import type { ImageAdapter } from './adapter'
11
11
  import type { ImageGenerationResult, StreamChunk } from '../../types'
@@ -5,7 +5,7 @@
5
5
  * This is a self-contained module with implementation, types, and JSDoc.
6
6
  */
7
7
 
8
- import { aiEventClient } from '../../event-client.js'
8
+ import { aiEventClient } from '@tanstack/ai-event-client'
9
9
  import { streamGenerationResult } from '../stream-generation-result.js'
10
10
  import type { TTSAdapter } from './adapter'
11
11
  import type { StreamChunk, TTSResult } from '../../types'
@@ -5,7 +5,7 @@
5
5
  * This is a self-contained module with implementation, types, and JSDoc.
6
6
  */
7
7
 
8
- import { aiEventClient } from '../../event-client.js'
8
+ import { aiEventClient } from '@tanstack/ai-event-client'
9
9
  import { streamGenerationResult } from '../stream-generation-result.js'
10
10
  import type { TranscriptionAdapter } from './adapter'
11
11
  import type { StreamChunk, TranscriptionResult } from '../../types'
@@ -7,7 +7,7 @@
7
7
  * @experimental Video generation is an experimental feature and may change.
8
8
  */
9
9
 
10
- import { aiEventClient } from '../../event-client.js'
10
+ import { aiEventClient } from '@tanstack/ai-event-client'
11
11
  import type { VideoAdapter } from './adapter'
12
12
  import type {
13
13
  StreamChunk,
@@ -5,7 +5,7 @@
5
5
  * This is a self-contained module with implementation, types, and JSDoc.
6
6
  */
7
7
 
8
- import { aiEventClient } from '../../event-client.js'
8
+ import { aiEventClient } from '@tanstack/ai-event-client'
9
9
  import { streamGenerationResult } from '../stream-generation-result.js'
10
10
  import type { SummarizeAdapter } from './adapter'
11
11
  import type {
package/src/index.ts CHANGED
@@ -70,14 +70,53 @@ export {
70
70
  combineStrategies,
71
71
  } from './activities/chat/agent-loop-strategies'
72
72
 
73
+ // Chat middleware
74
+ export type {
75
+ ChatMiddleware,
76
+ ChatMiddlewareContext,
77
+ ChatMiddlewarePhase,
78
+ ChatMiddlewareConfig,
79
+ ToolCallHookContext,
80
+ BeforeToolCallDecision,
81
+ AfterToolCallInfo,
82
+ IterationInfo,
83
+ ToolPhaseCompleteInfo,
84
+ UsageInfo,
85
+ FinishInfo,
86
+ AbortInfo,
87
+ ErrorInfo,
88
+ } from './activities/chat/middleware/index'
89
+
73
90
  // All types
74
91
  export * from './types'
75
92
 
76
93
  // Utility functions
77
94
  export { detectImageMimeType } from './utils'
78
95
 
79
- // Event client + event types
80
- export * from './event-client'
96
+ // Realtime
97
+ export { realtimeToken } from './realtime/index'
98
+ export type {
99
+ RealtimeToken,
100
+ RealtimeTokenAdapter,
101
+ RealtimeTokenOptions,
102
+ RealtimeSessionConfig,
103
+ VADConfig,
104
+ RealtimeMessage,
105
+ RealtimeMessagePart,
106
+ RealtimeTextPart,
107
+ RealtimeAudioPart,
108
+ RealtimeToolCallPart,
109
+ RealtimeToolResultPart,
110
+ RealtimeImagePart,
111
+ RealtimeStatus,
112
+ RealtimeMode,
113
+ AudioVisualization,
114
+ RealtimeEvent,
115
+ RealtimeEventPayloads,
116
+ RealtimeEventHandler,
117
+ RealtimeErrorCode,
118
+ RealtimeError,
119
+ } from './realtime/index'
81
120
 
82
121
  // Message converters
83
122
  export {