@tanstack/ai 0.6.3 → 0.8.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 (68) hide show
  1. package/README.md +6 -6
  2. package/dist/esm/activities/chat/index.d.ts +20 -0
  3. package/dist/esm/activities/chat/index.js +248 -213
  4. package/dist/esm/activities/chat/index.js.map +1 -1
  5. package/dist/esm/activities/chat/middleware/compose.d.ts +66 -0
  6. package/dist/esm/activities/chat/middleware/compose.js +327 -0
  7. package/dist/esm/activities/chat/middleware/compose.js.map +1 -0
  8. package/dist/esm/activities/chat/middleware/index.d.ts +2 -0
  9. package/dist/esm/activities/chat/middleware/tool-cache-middleware.d.ts +89 -0
  10. package/dist/esm/activities/chat/middleware/tool-cache-middleware.js +76 -0
  11. package/dist/esm/activities/chat/middleware/tool-cache-middleware.js.map +1 -0
  12. package/dist/esm/activities/chat/middleware/types.d.ts +307 -0
  13. package/dist/esm/activities/chat/stream/processor.d.ts +64 -40
  14. package/dist/esm/activities/chat/stream/processor.js +466 -218
  15. package/dist/esm/activities/chat/stream/processor.js.map +1 -1
  16. package/dist/esm/activities/chat/stream/types.d.ts +17 -0
  17. package/dist/esm/activities/chat/tools/tool-calls.d.ts +16 -1
  18. package/dist/esm/activities/chat/tools/tool-calls.js +148 -64
  19. package/dist/esm/activities/chat/tools/tool-calls.js.map +1 -1
  20. package/dist/esm/activities/generateImage/index.js +1 -1
  21. package/dist/esm/activities/generateImage/index.js.map +1 -1
  22. package/dist/esm/activities/generateSpeech/index.js +1 -1
  23. package/dist/esm/activities/generateSpeech/index.js.map +1 -1
  24. package/dist/esm/activities/generateTranscription/index.js +1 -1
  25. package/dist/esm/activities/generateTranscription/index.js.map +1 -1
  26. package/dist/esm/activities/generateVideo/index.js +1 -1
  27. package/dist/esm/activities/generateVideo/index.js.map +1 -1
  28. package/dist/esm/activities/summarize/index.js +1 -1
  29. package/dist/esm/activities/summarize/index.js.map +1 -1
  30. package/dist/esm/index.d.ts +3 -1
  31. package/dist/esm/index.js +2 -2
  32. package/dist/esm/middlewares/content-guard.d.ts +77 -0
  33. package/dist/esm/middlewares/content-guard.js +155 -0
  34. package/dist/esm/middlewares/content-guard.js.map +1 -0
  35. package/dist/esm/middlewares/index.d.ts +2 -0
  36. package/dist/esm/middlewares/index.js +7 -0
  37. package/dist/esm/middlewares/index.js.map +1 -0
  38. package/dist/esm/middlewares/tool-cache.d.ts +1 -0
  39. package/dist/esm/realtime/index.d.ts +30 -0
  40. package/dist/esm/realtime/index.js +8 -0
  41. package/dist/esm/realtime/index.js.map +1 -0
  42. package/dist/esm/realtime/types.d.ts +234 -0
  43. package/dist/esm/types.d.ts +18 -4
  44. package/package.json +6 -6
  45. package/src/activities/chat/index.ts +322 -256
  46. package/src/activities/chat/middleware/compose.ts +392 -0
  47. package/src/activities/chat/middleware/index.ts +17 -0
  48. package/src/activities/chat/middleware/tool-cache-middleware.ts +189 -0
  49. package/src/activities/chat/middleware/types.ts +419 -0
  50. package/src/activities/chat/stream/processor.ts +630 -259
  51. package/src/activities/chat/stream/types.ts +18 -0
  52. package/src/activities/chat/tools/tool-calls.ts +225 -87
  53. package/src/activities/generateImage/index.ts +1 -1
  54. package/src/activities/generateSpeech/index.ts +1 -1
  55. package/src/activities/generateTranscription/index.ts +1 -1
  56. package/src/activities/generateVideo/index.ts +1 -1
  57. package/src/activities/summarize/index.ts +1 -1
  58. package/src/index.ts +41 -2
  59. package/src/middlewares/content-guard.ts +285 -0
  60. package/src/middlewares/index.ts +13 -0
  61. package/src/middlewares/tool-cache.ts +6 -0
  62. package/src/realtime/index.ts +38 -0
  63. package/src/realtime/types.ts +294 -0
  64. package/src/types.ts +19 -2
  65. package/dist/esm/event-client.d.ts +0 -394
  66. package/dist/esm/event-client.js +0 -13
  67. package/dist/esm/event-client.js.map +0 -1
  68. package/src/event-client.ts +0 -497
@@ -5,9 +5,13 @@
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 { devtoolsMiddleware } from '@tanstack/ai-event-client'
9
9
  import { streamToText } from '../../stream-to-response.js'
10
- import { ToolCallManager, executeToolCalls } from './tools/tool-calls'
10
+ import {
11
+ MiddlewareAbortError,
12
+ ToolCallManager,
13
+ executeToolCalls,
14
+ } from './tools/tool-calls'
11
15
  import {
12
16
  convertSchemaToJsonSchema,
13
17
  isStandardSchema,
@@ -15,6 +19,7 @@ import {
15
19
  } from './tools/schema-converter'
16
20
  import { maxIterations as maxIterationsStrategy } from './agent-loop-strategies'
17
21
  import { convertMessagesToModelMessages } from './messages'
22
+ import { MiddlewareRunner } from './middleware/compose'
18
23
  import type {
19
24
  ApprovalRequest,
20
25
  ClientToolRequest,
@@ -38,6 +43,12 @@ import type {
38
43
  ToolCallEndEvent,
39
44
  ToolCallStartEvent,
40
45
  } from '../../types'
46
+ import type {
47
+ ChatMiddleware,
48
+ ChatMiddlewareConfig,
49
+ ChatMiddlewareContext,
50
+ ChatMiddlewarePhase,
51
+ } from './middleware/types'
41
52
 
42
53
  // ===========================
43
54
  // Activity Kind
@@ -132,6 +143,25 @@ export interface TextActivityOptions<
132
143
  * ```
133
144
  */
134
145
  stream?: TStream
146
+ /**
147
+ * Optional middleware array for observing/transforming chat behavior.
148
+ * Middleware hooks are called in array order. See {@link ChatMiddleware} for available hooks.
149
+ *
150
+ * @example
151
+ * ```ts
152
+ * const stream = chat({
153
+ * adapter: openaiText('gpt-4o'),
154
+ * messages: [...],
155
+ * middleware: [loggingMiddleware, redactionMiddleware],
156
+ * })
157
+ * ```
158
+ */
159
+ middleware?: Array<ChatMiddleware>
160
+ /**
161
+ * Opaque user-provided context value passed to middleware hooks.
162
+ * Can be used to pass request-scoped data (e.g., user ID, request context).
163
+ */
164
+ context?: unknown
135
165
  }
136
166
 
137
167
  // ===========================
@@ -191,6 +221,8 @@ interface TextEngineConfig<
191
221
  adapter: TAdapter
192
222
  systemPrompts?: Array<string>
193
223
  params: TParams
224
+ middleware?: Array<ChatMiddleware>
225
+ context?: unknown
194
226
  }
195
227
 
196
228
  type ToolPhaseResult = 'continue' | 'stop' | 'wait'
@@ -201,9 +233,9 @@ class TextEngine<
201
233
  TParams extends TextOptions<any, any> = TextOptions<any>,
202
234
  > {
203
235
  private readonly adapter: TAdapter
204
- private readonly params: TParams
205
- private readonly systemPrompts: Array<string>
206
- private readonly tools: ReadonlyArray<Tool>
236
+ private params: TParams
237
+ private systemPrompts: Array<string>
238
+ private tools: Array<Tool>
207
239
  private readonly loopStrategy: AgentLoopStrategy
208
240
  private readonly toolCallManager: ToolCallManager
209
241
  private readonly initialMessageCount: number
@@ -222,7 +254,6 @@ class TextEngine<
222
254
  private eventOptions?: Record<string, unknown>
223
255
  private eventToolNames?: Array<string>
224
256
  private finishedEvent: RunFinishedEvent | null = null
225
- private shouldEmitStreamEnd = true
226
257
  private earlyTermination = false
227
258
  private toolPhase: ToolPhaseResult = 'continue'
228
259
  private cyclePhase: CyclePhase = 'processText'
@@ -230,6 +261,14 @@ class TextEngine<
230
261
  private readonly initialApprovals: Map<string, boolean>
231
262
  private readonly initialClientToolResults: Map<string, any>
232
263
 
264
+ // Middleware support
265
+ private readonly middlewareRunner: MiddlewareRunner
266
+ private readonly middlewareCtx: ChatMiddlewareContext
267
+ private readonly deferredPromises: Array<Promise<unknown>> = []
268
+ private abortReason?: string
269
+ private middlewareAbortController?: AbortController
270
+ private terminalHookCalled = false
271
+
233
272
  constructor(config: TextEngineConfig<TAdapter, TParams>) {
234
273
  this.adapter = config.adapter
235
274
  this.params = config.params
@@ -260,6 +299,47 @@ class TextEngine<
260
299
  ? { signal: config.params.abortController.signal }
261
300
  : undefined
262
301
  this.effectiveSignal = config.params.abortController?.signal
302
+
303
+ // Initialize middleware — devtools middleware is always first
304
+ const allMiddleware = [devtoolsMiddleware(), ...(config.middleware || [])]
305
+ this.middlewareRunner = new MiddlewareRunner(allMiddleware)
306
+ this.middlewareAbortController = new AbortController()
307
+ this.middlewareCtx = {
308
+ requestId: this.requestId,
309
+ streamId: this.streamId,
310
+ conversationId: config.params.conversationId,
311
+ phase: 'init' as ChatMiddlewarePhase,
312
+ iteration: 0,
313
+ chunkIndex: 0,
314
+ signal: this.effectiveSignal,
315
+ abort: (reason?: string) => {
316
+ this.abortReason = reason
317
+ this.middlewareAbortController?.abort(reason)
318
+ },
319
+ context: config.context,
320
+ defer: (promise: Promise<unknown>) => {
321
+ this.deferredPromises.push(promise)
322
+ },
323
+ // Provider / adapter info
324
+ provider: config.adapter.name,
325
+ model: config.params.model,
326
+ source: 'server',
327
+ streaming: true,
328
+ // Config-derived (updated in beforeRun and applyMiddlewareConfig)
329
+ systemPrompts: this.systemPrompts,
330
+ toolNames: undefined,
331
+ options: undefined,
332
+ modelOptions: config.params.modelOptions,
333
+ // Computed
334
+ messageCount: this.initialMessageCount,
335
+ hasTools: this.tools.length > 0,
336
+ // Mutable per-iteration
337
+ currentMessageId: null,
338
+ accumulatedContent: '',
339
+ // References
340
+ messages: this.messages,
341
+ createId: (prefix: string) => this.createId(prefix),
342
+ }
263
343
  }
264
344
 
265
345
  /** Get the accumulated content after the chat loop completes */
@@ -276,19 +356,41 @@ class TextEngine<
276
356
  this.beforeRun()
277
357
 
278
358
  try {
359
+ // Run initial onConfig (phase = init)
360
+ this.middlewareCtx.phase = 'init'
361
+ const initialConfig = this.buildMiddlewareConfig()
362
+ const transformedConfig = await this.middlewareRunner.runOnConfig(
363
+ this.middlewareCtx,
364
+ initialConfig,
365
+ )
366
+ this.applyMiddlewareConfig(transformedConfig)
367
+
368
+ // Run onStart (devtools middleware emits text:request:started and initial messages here)
369
+ await this.middlewareRunner.runOnStart(this.middlewareCtx)
370
+
279
371
  const pendingPhase = yield* this.checkForPendingToolCalls()
280
372
  if (pendingPhase === 'wait') {
281
373
  return
282
374
  }
283
375
 
284
376
  do {
285
- if (this.earlyTermination || this.isAborted()) {
377
+ if (this.earlyTermination || this.isCancelled()) {
286
378
  return
287
379
  }
288
380
 
289
- this.beginCycle()
381
+ await this.beginCycle()
290
382
 
291
383
  if (this.cyclePhase === 'processText') {
384
+ // Run onConfig before each model call (phase = beforeModel)
385
+ this.middlewareCtx.phase = 'beforeModel'
386
+ this.middlewareCtx.iteration = this.iterationCount
387
+ const iterConfig = this.buildMiddlewareConfig()
388
+ const transformedConfig = await this.middlewareRunner.runOnConfig(
389
+ this.middlewareCtx,
390
+ iterConfig,
391
+ )
392
+ this.applyMiddlewareConfig(transformedConfig)
393
+
292
394
  yield* this.streamModelResponse()
293
395
  } else {
294
396
  yield* this.processToolCalls()
@@ -296,8 +398,53 @@ class TextEngine<
296
398
 
297
399
  this.endCycle()
298
400
  } while (this.shouldContinue())
401
+
402
+ // Call terminal onFinish hook (skip when waiting for client — stream is paused, not finished)
403
+ if (!this.terminalHookCalled && this.toolPhase !== 'wait') {
404
+ this.terminalHookCalled = true
405
+ await this.middlewareRunner.runOnFinish(this.middlewareCtx, {
406
+ finishReason: this.lastFinishReason,
407
+ duration: Date.now() - this.streamStartTime,
408
+ content: this.accumulatedContent,
409
+ usage: this.finishedEvent?.usage,
410
+ })
411
+ }
412
+ } catch (error: unknown) {
413
+ if (!this.terminalHookCalled) {
414
+ this.terminalHookCalled = true
415
+ if (error instanceof MiddlewareAbortError) {
416
+ // Middleware abort decision — call onAbort, not onError
417
+ this.abortReason = error.message
418
+ await this.middlewareRunner.runOnAbort(this.middlewareCtx, {
419
+ reason: error.message,
420
+ duration: Date.now() - this.streamStartTime,
421
+ })
422
+ } else {
423
+ // Genuine error — call onError
424
+ await this.middlewareRunner.runOnError(this.middlewareCtx, {
425
+ error,
426
+ duration: Date.now() - this.streamStartTime,
427
+ })
428
+ }
429
+ }
430
+ // Don't rethrow middleware abort errors — the run just stops gracefully
431
+ if (!(error instanceof MiddlewareAbortError)) {
432
+ throw error
433
+ }
299
434
  } finally {
300
- this.afterRun()
435
+ // Check for abort terminal hook
436
+ if (!this.terminalHookCalled && this.isCancelled()) {
437
+ this.terminalHookCalled = true
438
+ await this.middlewareRunner.runOnAbort(this.middlewareCtx, {
439
+ reason: this.abortReason,
440
+ duration: Date.now() - this.streamStartTime,
441
+ })
442
+ }
443
+
444
+ // Await deferred promises (non-blocking side effects)
445
+ if (this.deferredPromises.length > 0) {
446
+ await Promise.allSettled(this.deferredPromises)
447
+ }
301
448
  }
302
449
  }
303
450
 
@@ -305,7 +452,7 @@ class TextEngine<
305
452
  this.streamStartTime = Date.now()
306
453
  const { tools, temperature, topP, maxTokens, metadata } = this.params
307
454
 
308
- // Gather flattened options into an object for event emission
455
+ // Gather flattened options into an object for context
309
456
  const options: Record<string, unknown> = {}
310
457
  if (temperature !== undefined) options.temperature = temperature
311
458
  if (topP !== undefined) options.topP = topP
@@ -315,70 +462,14 @@ class TextEngine<
315
462
  this.eventOptions = Object.keys(options).length > 0 ? options : undefined
316
463
  this.eventToolNames = tools?.map((t) => t.name)
317
464
 
318
- aiEventClient.emit('text:request:started', {
319
- ...this.buildTextEventContext(),
320
- timestamp: Date.now(),
321
- })
322
-
323
- // Always emit messages for tracking:
324
- // - For existing conversations (with conversationId): only emit the latest user message
325
- // - For new conversations (without conversationId): emit all messages for reconstruction
326
- const messagesToEmit = this.params.conversationId
327
- ? this.messages.slice(-1).filter((m) => m.role === 'user')
328
- : this.messages
329
-
330
- messagesToEmit.forEach((message, index) => {
331
- const messageIndex = this.params.conversationId
332
- ? this.messages.length - 1
333
- : index
334
- const messageId = this.createId('msg')
335
- const baseContext = this.buildTextEventContext()
336
- const content = this.getContentString(message.content)
337
-
338
- aiEventClient.emit('text:message:created', {
339
- ...baseContext,
340
- messageId,
341
- role: message.role,
342
- content,
343
- toolCalls: message.toolCalls,
344
- messageIndex,
345
- timestamp: Date.now(),
346
- })
347
-
348
- if (message.role === 'user') {
349
- aiEventClient.emit('text:message:user', {
350
- ...baseContext,
351
- messageId,
352
- role: 'user',
353
- content,
354
- messageIndex,
355
- timestamp: Date.now(),
356
- })
357
- }
358
- })
359
- }
360
-
361
- private afterRun(): void {
362
- if (!this.shouldEmitStreamEnd) {
363
- return
364
- }
365
-
366
- const now = Date.now()
367
- // Emit text:request:completed with final state
368
- aiEventClient.emit('text:request:completed', {
369
- ...this.buildTextEventContext(),
370
- content: this.accumulatedContent,
371
- messageId: this.currentMessageId || undefined,
372
- finishReason: this.lastFinishReason || undefined,
373
- usage: this.finishedEvent?.usage,
374
- duration: now - this.streamStartTime,
375
- timestamp: now,
376
- })
465
+ // Update middleware context with computed fields
466
+ this.middlewareCtx.options = this.eventOptions
467
+ this.middlewareCtx.toolNames = this.eventToolNames
377
468
  }
378
469
 
379
- private beginCycle(): void {
470
+ private async beginCycle(): Promise<void> {
380
471
  if (this.cyclePhase === 'processText') {
381
- this.beginIteration()
472
+ await this.beginIteration()
382
473
  }
383
474
  }
384
475
 
@@ -392,27 +483,28 @@ class TextEngine<
392
483
  this.iterationCount++
393
484
  }
394
485
 
395
- private beginIteration(): void {
486
+ private async beginIteration(): Promise<void> {
396
487
  this.currentMessageId = this.createId('msg')
397
488
  this.accumulatedContent = ''
398
489
  this.finishedEvent = null
399
490
 
400
- const baseContext = this.buildTextEventContext()
401
- aiEventClient.emit('text:message:created', {
402
- ...baseContext,
491
+ // Update mutable context fields
492
+ this.middlewareCtx.currentMessageId = this.currentMessageId
493
+ this.middlewareCtx.accumulatedContent = ''
494
+
495
+ // Notify middleware of new iteration (devtools emits assistant message:created here)
496
+ await this.middlewareRunner.runOnIteration(this.middlewareCtx, {
497
+ iteration: this.iterationCount,
403
498
  messageId: this.currentMessageId,
404
- role: 'assistant',
405
- content: '',
406
- timestamp: Date.now(),
407
499
  })
408
500
  }
409
501
 
410
502
  private async *streamModelResponse(): AsyncGenerator<StreamChunk> {
411
503
  const { temperature, topP, maxTokens, metadata, modelOptions } = this.params
412
- const tools = this.params.tools
504
+ const tools = this.tools
413
505
 
414
506
  // Convert tool schemas to JSON Schema before passing to adapter
415
- const toolsWithJsonSchemas = tools?.map((tool) => ({
507
+ const toolsWithJsonSchemas = tools.map((tool) => ({
416
508
  ...tool,
417
509
  inputSchema: tool.inputSchema
418
510
  ? convertSchemaToJsonSchema(tool.inputSchema)
@@ -422,6 +514,8 @@ class TextEngine<
422
514
  : undefined,
423
515
  }))
424
516
 
517
+ this.middlewareCtx.phase = 'modelStream'
518
+
425
519
  for await (const chunk of this.adapter.chatStream({
426
520
  model: this.params.model,
427
521
  messages: this.messages,
@@ -434,14 +528,27 @@ class TextEngine<
434
528
  modelOptions,
435
529
  systemPrompts: this.systemPrompts,
436
530
  })) {
437
- if (this.isAborted()) {
531
+ if (this.isCancelled()) {
438
532
  break
439
533
  }
440
534
 
441
535
  this.totalChunkCount++
442
536
 
443
- yield chunk
444
- this.handleStreamChunk(chunk)
537
+ // Pipe chunk through middleware (devtools middleware observes and emits events)
538
+ const outputChunks = await this.middlewareRunner.runOnChunk(
539
+ this.middlewareCtx,
540
+ chunk,
541
+ )
542
+ for (const outputChunk of outputChunks) {
543
+ yield outputChunk
544
+ this.handleStreamChunk(outputChunk)
545
+ this.middlewareCtx.chunkIndex++
546
+ }
547
+
548
+ // Handle usage via middleware
549
+ if (chunk.type === 'RUN_FINISHED' && chunk.usage) {
550
+ await this.middlewareRunner.runOnUsage(this.middlewareCtx, chunk.usage)
551
+ }
445
552
 
446
553
  if (this.earlyTermination) {
447
554
  break
@@ -492,100 +599,36 @@ class TextEngine<
492
599
  } else {
493
600
  this.accumulatedContent += chunk.delta
494
601
  }
495
- aiEventClient.emit('text:chunk:content', {
496
- ...this.buildTextEventContext(),
497
- messageId: this.currentMessageId || undefined,
498
- content: this.accumulatedContent,
499
- delta: chunk.delta,
500
- timestamp: Date.now(),
501
- })
602
+ this.middlewareCtx.accumulatedContent = this.accumulatedContent
502
603
  }
503
604
 
504
605
  private handleToolCallStartEvent(chunk: ToolCallStartEvent): void {
505
606
  this.toolCallManager.addToolCallStartEvent(chunk)
506
- aiEventClient.emit('text:chunk:tool-call', {
507
- ...this.buildTextEventContext(),
508
- messageId: this.currentMessageId || undefined,
509
- toolCallId: chunk.toolCallId,
510
- toolName: chunk.toolName,
511
- index: chunk.index ?? 0,
512
- arguments: '',
513
- timestamp: Date.now(),
514
- })
515
607
  }
516
608
 
517
609
  private handleToolCallArgsEvent(chunk: ToolCallArgsEvent): void {
518
610
  this.toolCallManager.addToolCallArgsEvent(chunk)
519
- aiEventClient.emit('text:chunk:tool-call', {
520
- ...this.buildTextEventContext(),
521
- messageId: this.currentMessageId || undefined,
522
- toolCallId: chunk.toolCallId,
523
- toolName: '',
524
- index: 0,
525
- arguments: chunk.delta,
526
- timestamp: Date.now(),
527
- })
528
611
  }
529
612
 
530
613
  private handleToolCallEndEvent(chunk: ToolCallEndEvent): void {
531
614
  this.toolCallManager.completeToolCall(chunk)
532
- aiEventClient.emit('text:chunk:tool-result', {
533
- ...this.buildTextEventContext(),
534
- messageId: this.currentMessageId || undefined,
535
- toolCallId: chunk.toolCallId,
536
- result: chunk.result || '',
537
- timestamp: Date.now(),
538
- })
539
615
  }
540
616
 
541
617
  private handleRunFinishedEvent(chunk: RunFinishedEvent): void {
542
- aiEventClient.emit('text:chunk:done', {
543
- ...this.buildTextEventContext(),
544
- messageId: this.currentMessageId || undefined,
545
- finishReason: chunk.finishReason,
546
- usage: chunk.usage,
547
- timestamp: Date.now(),
548
- })
549
-
550
- if (chunk.usage) {
551
- aiEventClient.emit('text:usage', {
552
- ...this.buildTextEventContext(),
553
- messageId: this.currentMessageId || undefined,
554
- usage: chunk.usage,
555
- timestamp: Date.now(),
556
- })
557
- }
558
-
559
618
  this.finishedEvent = chunk
560
619
  this.lastFinishReason = chunk.finishReason
561
620
  }
562
621
 
563
622
  private handleRunErrorEvent(
564
- chunk: Extract<StreamChunk, { type: 'RUN_ERROR' }>,
623
+ _chunk: Extract<StreamChunk, { type: 'RUN_ERROR' }>,
565
624
  ): void {
566
- aiEventClient.emit('text:chunk:error', {
567
- ...this.buildTextEventContext(),
568
- messageId: this.currentMessageId || undefined,
569
- error: chunk.error.message,
570
- timestamp: Date.now(),
571
- })
572
625
  this.earlyTermination = true
573
- this.shouldEmitStreamEnd = false
574
626
  }
575
627
 
576
628
  private handleStepFinishedEvent(
577
- chunk: Extract<StreamChunk, { type: 'STEP_FINISHED' }>,
629
+ _chunk: Extract<StreamChunk, { type: 'STEP_FINISHED' }>,
578
630
  ): void {
579
- // Handle thinking/reasoning content from STEP_FINISHED events
580
- if (chunk.content || chunk.delta) {
581
- aiEventClient.emit('text:chunk:thinking', {
582
- ...this.buildTextEventContext(),
583
- messageId: this.currentMessageId || undefined,
584
- content: chunk.content || '',
585
- delta: chunk.delta,
586
- timestamp: Date.now(),
587
- })
588
- }
631
+ // State tracking for STEP_FINISHED is handled by middleware
589
632
  }
590
633
 
591
634
  private async *checkForPendingToolCalls(): AsyncGenerator<
@@ -608,17 +651,52 @@ class TextEngine<
608
651
  approvals,
609
652
  clientToolResults,
610
653
  (eventName, data) => this.createCustomEventChunk(eventName, data),
654
+ {
655
+ onBeforeToolCall: async (toolCall, tool, args) => {
656
+ const hookCtx = {
657
+ toolCall,
658
+ tool,
659
+ args,
660
+ toolName: toolCall.function.name,
661
+ toolCallId: toolCall.id,
662
+ }
663
+ return this.middlewareRunner.runOnBeforeToolCall(
664
+ this.middlewareCtx,
665
+ hookCtx,
666
+ )
667
+ },
668
+ onAfterToolCall: async (info) => {
669
+ await this.middlewareRunner.runOnAfterToolCall(
670
+ this.middlewareCtx,
671
+ info,
672
+ )
673
+ },
674
+ },
611
675
  )
612
676
 
613
677
  // Consume the async generator, yielding custom events and collecting the return value
614
678
  const executionResult = yield* this.drainToolCallGenerator(generator)
615
679
 
680
+ // Check if middleware aborted during pending tool execution
681
+ if (this.isMiddlewareAborted()) {
682
+ this.setToolPhase('stop')
683
+ return 'stop'
684
+ }
685
+
686
+ // Notify middleware of tool phase completion (devtools emits aggregate events here)
687
+ await this.middlewareRunner.runOnToolPhaseComplete(this.middlewareCtx, {
688
+ toolCalls: pendingToolCalls,
689
+ results: executionResult.results,
690
+ needsApproval: executionResult.needsApproval,
691
+ needsClientExecution: executionResult.needsClientExecution,
692
+ })
693
+
616
694
  if (
617
695
  executionResult.needsApproval.length > 0 ||
618
696
  executionResult.needsClientExecution.length > 0
619
697
  ) {
620
698
  if (executionResult.results.length > 0) {
621
- for (const chunk of this.emitToolResults(
699
+ for (const chunk of this.buildToolResultChunks(
622
700
  executionResult.results,
623
701
  finishEvent,
624
702
  )) {
@@ -626,25 +704,25 @@ class TextEngine<
626
704
  }
627
705
  }
628
706
 
629
- for (const chunk of this.emitApprovalRequests(
707
+ for (const chunk of this.buildApprovalChunks(
630
708
  executionResult.needsApproval,
631
709
  finishEvent,
632
710
  )) {
633
711
  yield chunk
634
712
  }
635
713
 
636
- for (const chunk of this.emitClientToolInputs(
714
+ for (const chunk of this.buildClientToolChunks(
637
715
  executionResult.needsClientExecution,
638
716
  finishEvent,
639
717
  )) {
640
718
  yield chunk
641
719
  }
642
720
 
643
- this.shouldEmitStreamEnd = false
721
+ this.setToolPhase('wait')
644
722
  return 'wait'
645
723
  }
646
724
 
647
- const toolResultChunks = this.emitToolResults(
725
+ const toolResultChunks = this.buildToolResultChunks(
648
726
  executionResult.results,
649
727
  finishEvent,
650
728
  )
@@ -672,6 +750,8 @@ class TextEngine<
672
750
 
673
751
  this.addAssistantToolCallMessage(toolCalls)
674
752
 
753
+ this.middlewareCtx.phase = 'beforeTools'
754
+
675
755
  const { approvals, clientToolResults } = this.collectClientState()
676
756
 
677
757
  const generator = executeToolCalls(
@@ -680,17 +760,54 @@ class TextEngine<
680
760
  approvals,
681
761
  clientToolResults,
682
762
  (eventName, data) => this.createCustomEventChunk(eventName, data),
763
+ {
764
+ onBeforeToolCall: async (toolCall, tool, args) => {
765
+ const hookCtx = {
766
+ toolCall,
767
+ tool,
768
+ args,
769
+ toolName: toolCall.function.name,
770
+ toolCallId: toolCall.id,
771
+ }
772
+ return this.middlewareRunner.runOnBeforeToolCall(
773
+ this.middlewareCtx,
774
+ hookCtx,
775
+ )
776
+ },
777
+ onAfterToolCall: async (info) => {
778
+ await this.middlewareRunner.runOnAfterToolCall(
779
+ this.middlewareCtx,
780
+ info,
781
+ )
782
+ },
783
+ },
683
784
  )
684
785
 
685
786
  // Consume the async generator, yielding custom events and collecting the return value
686
787
  const executionResult = yield* this.drainToolCallGenerator(generator)
687
788
 
789
+ this.middlewareCtx.phase = 'afterTools'
790
+
791
+ // Check if middleware aborted during tool execution
792
+ if (this.isMiddlewareAborted()) {
793
+ this.setToolPhase('stop')
794
+ return
795
+ }
796
+
797
+ // Notify middleware of tool phase completion (devtools emits aggregate events here)
798
+ await this.middlewareRunner.runOnToolPhaseComplete(this.middlewareCtx, {
799
+ toolCalls,
800
+ results: executionResult.results,
801
+ needsApproval: executionResult.needsApproval,
802
+ needsClientExecution: executionResult.needsClientExecution,
803
+ })
804
+
688
805
  if (
689
806
  executionResult.needsApproval.length > 0 ||
690
807
  executionResult.needsClientExecution.length > 0
691
808
  ) {
692
809
  if (executionResult.results.length > 0) {
693
- for (const chunk of this.emitToolResults(
810
+ for (const chunk of this.buildToolResultChunks(
694
811
  executionResult.results,
695
812
  finishEvent,
696
813
  )) {
@@ -698,14 +815,14 @@ class TextEngine<
698
815
  }
699
816
  }
700
817
 
701
- for (const chunk of this.emitApprovalRequests(
818
+ for (const chunk of this.buildApprovalChunks(
702
819
  executionResult.needsApproval,
703
820
  finishEvent,
704
821
  )) {
705
822
  yield chunk
706
823
  }
707
824
 
708
- for (const chunk of this.emitClientToolInputs(
825
+ for (const chunk of this.buildClientToolChunks(
709
826
  executionResult.needsClientExecution,
710
827
  finishEvent,
711
828
  )) {
@@ -716,7 +833,7 @@ class TextEngine<
716
833
  return
717
834
  }
718
835
 
719
- const toolResultChunks = this.emitToolResults(
836
+ const toolResultChunks = this.buildToolResultChunks(
720
837
  executionResult.results,
721
838
  finishEvent,
722
839
  )
@@ -739,7 +856,6 @@ class TextEngine<
739
856
  }
740
857
 
741
858
  private addAssistantToolCallMessage(toolCalls: Array<ToolCall>): void {
742
- const messageId = this.currentMessageId ?? this.createId('msg')
743
859
  this.messages = [
744
860
  ...this.messages,
745
861
  {
@@ -748,15 +864,6 @@ class TextEngine<
748
864
  toolCalls,
749
865
  },
750
866
  ]
751
-
752
- aiEventClient.emit('text:message:created', {
753
- ...this.buildTextEventContext(),
754
- messageId,
755
- role: 'assistant',
756
- content: this.accumulatedContent || '',
757
- toolCalls,
758
- timestamp: Date.now(),
759
- })
760
867
  }
761
868
 
762
869
  /**
@@ -837,24 +944,13 @@ class TextEngine<
837
944
  return { approvals, clientToolResults }
838
945
  }
839
946
 
840
- private emitApprovalRequests(
947
+ private buildApprovalChunks(
841
948
  approvals: Array<ApprovalRequest>,
842
949
  finishEvent: RunFinishedEvent,
843
950
  ): Array<StreamChunk> {
844
951
  const chunks: Array<StreamChunk> = []
845
952
 
846
953
  for (const approval of approvals) {
847
- aiEventClient.emit('tools:approval:requested', {
848
- ...this.buildTextEventContext(),
849
- messageId: this.currentMessageId || undefined,
850
- toolCallId: approval.toolCallId,
851
- toolName: approval.toolName,
852
- input: approval.input,
853
- approvalId: approval.approvalId,
854
- timestamp: Date.now(),
855
- })
856
-
857
- // Emit a CUSTOM event for approval requests
858
954
  chunks.push({
859
955
  type: 'CUSTOM',
860
956
  timestamp: Date.now(),
@@ -875,23 +971,13 @@ class TextEngine<
875
971
  return chunks
876
972
  }
877
973
 
878
- private emitClientToolInputs(
974
+ private buildClientToolChunks(
879
975
  clientRequests: Array<ClientToolRequest>,
880
976
  finishEvent: RunFinishedEvent,
881
977
  ): Array<StreamChunk> {
882
978
  const chunks: Array<StreamChunk> = []
883
979
 
884
980
  for (const clientTool of clientRequests) {
885
- aiEventClient.emit('tools:input:available', {
886
- ...this.buildTextEventContext(),
887
- messageId: this.currentMessageId || undefined,
888
- toolCallId: clientTool.toolCallId,
889
- toolName: clientTool.toolName,
890
- input: clientTool.input,
891
- timestamp: Date.now(),
892
- })
893
-
894
- // Emit a CUSTOM event for client tool inputs
895
981
  chunks.push({
896
982
  type: 'CUSTOM',
897
983
  timestamp: Date.now(),
@@ -908,26 +994,15 @@ class TextEngine<
908
994
  return chunks
909
995
  }
910
996
 
911
- private emitToolResults(
997
+ private buildToolResultChunks(
912
998
  results: Array<ToolResult>,
913
999
  finishEvent: RunFinishedEvent,
914
1000
  ): Array<StreamChunk> {
915
1001
  const chunks: Array<StreamChunk> = []
916
1002
 
917
1003
  for (const result of results) {
918
- aiEventClient.emit('tools:call:completed', {
919
- ...this.buildTextEventContext(),
920
- messageId: this.currentMessageId || undefined,
921
- toolCallId: result.toolCallId,
922
- toolName: result.toolName,
923
- result: result.result,
924
- duration: result.duration ?? 0,
925
- timestamp: Date.now(),
926
- })
927
-
928
1004
  const content = JSON.stringify(result.result)
929
1005
 
930
- // Emit TOOL_CALL_END event
931
1006
  chunks.push({
932
1007
  type: 'TOOL_CALL_END',
933
1008
  timestamp: Date.now(),
@@ -945,14 +1020,6 @@ class TextEngine<
945
1020
  toolCallId: result.toolCallId,
946
1021
  },
947
1022
  ]
948
-
949
- aiEventClient.emit('text:message:created', {
950
- ...this.buildTextEventContext(),
951
- messageId: this.createId('msg'),
952
- role: 'tool',
953
- content,
954
- timestamp: Date.now(),
955
- })
956
1023
  }
957
1024
 
958
1025
  return chunks
@@ -1028,55 +1095,50 @@ class TextEngine<
1028
1095
  return !!this.effectiveSignal?.aborted
1029
1096
  }
1030
1097
 
1031
- private buildTextEventContext(): {
1032
- requestId: string
1033
- streamId: string
1034
- provider: string
1035
- model: string
1036
- clientId?: string
1037
- source?: 'client' | 'server'
1038
- systemPrompts?: Array<string>
1039
- toolNames?: Array<string>
1040
- options?: Record<string, unknown>
1041
- modelOptions?: Record<string, unknown>
1042
- messageCount: number
1043
- hasTools: boolean
1044
- streaming: boolean
1045
- } {
1098
+ private isMiddlewareAborted(): boolean {
1099
+ return !!this.middlewareAbortController?.signal.aborted
1100
+ }
1101
+
1102
+ private isCancelled(): boolean {
1103
+ return this.isAborted() || this.isMiddlewareAborted()
1104
+ }
1105
+
1106
+ private buildMiddlewareConfig(): ChatMiddlewareConfig {
1046
1107
  return {
1047
- requestId: this.requestId,
1048
- streamId: this.streamId,
1049
- provider: this.adapter.name,
1050
- model: this.params.model,
1051
- clientId: this.params.conversationId,
1052
- source: 'server',
1053
- systemPrompts:
1054
- this.systemPrompts.length > 0 ? this.systemPrompts : undefined,
1055
- toolNames: this.eventToolNames,
1056
- options: this.eventOptions,
1057
- modelOptions: this.params.modelOptions as
1058
- | Record<string, unknown>
1059
- | undefined,
1060
- messageCount: this.initialMessageCount,
1061
- hasTools: this.tools.length > 0,
1062
- streaming: true,
1108
+ messages: this.messages,
1109
+ systemPrompts: [...this.systemPrompts],
1110
+ tools: [...this.tools],
1111
+ temperature: this.params.temperature,
1112
+ topP: this.params.topP,
1113
+ maxTokens: this.params.maxTokens,
1114
+ metadata: this.params.metadata,
1115
+ modelOptions: this.params.modelOptions,
1063
1116
  }
1064
1117
  }
1065
1118
 
1066
- private getContentString(content: ModelMessage['content']): string {
1067
- if (typeof content === 'string') return content
1068
- const text =
1069
- content
1070
- ?.map((part) => (part.type === 'text' ? part.content : ''))
1071
- .join('') || ''
1072
- return text
1119
+ private applyMiddlewareConfig(config: ChatMiddlewareConfig): void {
1120
+ this.messages = config.messages
1121
+ this.systemPrompts = config.systemPrompts
1122
+ this.tools = config.tools
1123
+ this.params = {
1124
+ ...this.params,
1125
+ temperature: config.temperature,
1126
+ topP: config.topP,
1127
+ maxTokens: config.maxTokens,
1128
+ metadata: config.metadata,
1129
+ modelOptions: config.modelOptions,
1130
+ }
1131
+
1132
+ // Sync context fields that depend on config
1133
+ this.middlewareCtx.messages = this.messages
1134
+ this.middlewareCtx.systemPrompts = this.systemPrompts
1135
+ this.middlewareCtx.hasTools = this.tools.length > 0
1136
+ this.middlewareCtx.toolNames = this.tools.map((t) => t.name)
1137
+ this.middlewareCtx.modelOptions = config.modelOptions
1073
1138
  }
1074
1139
 
1075
1140
  private setToolPhase(phase: ToolPhaseResult): void {
1076
1141
  this.toolPhase = phase
1077
- if (phase === 'wait') {
1078
- this.shouldEmitStreamEnd = false
1079
- }
1080
1142
  }
1081
1143
 
1082
1144
  /**
@@ -1236,7 +1298,7 @@ export function chat<
1236
1298
  async function* runStreamingText(
1237
1299
  options: TextActivityOptions<AnyTextAdapter, undefined, true>,
1238
1300
  ): AsyncIterable<StreamChunk> {
1239
- const { adapter, ...textOptions } = options
1301
+ const { adapter, middleware, context, ...textOptions } = options
1240
1302
  const model = adapter.model
1241
1303
 
1242
1304
  const engine = new TextEngine({
@@ -1245,6 +1307,8 @@ async function* runStreamingText(
1245
1307
  Record<string, any>,
1246
1308
  Record<string, any>
1247
1309
  >,
1310
+ middleware,
1311
+ context,
1248
1312
  })
1249
1313
 
1250
1314
  for await (const chunk of engine.run()) {
@@ -1276,7 +1340,7 @@ function runNonStreamingText(
1276
1340
  async function runAgenticStructuredOutput<TSchema extends SchemaInput>(
1277
1341
  options: TextActivityOptions<AnyTextAdapter, TSchema, boolean>,
1278
1342
  ): Promise<InferSchemaType<TSchema>> {
1279
- const { adapter, outputSchema, ...textOptions } = options
1343
+ const { adapter, outputSchema, middleware, context, ...textOptions } = options
1280
1344
  const model = adapter.model
1281
1345
 
1282
1346
  if (!outputSchema) {
@@ -1290,6 +1354,8 @@ async function runAgenticStructuredOutput<TSchema extends SchemaInput>(
1290
1354
  Record<string, unknown>,
1291
1355
  Record<string, unknown>
1292
1356
  >,
1357
+ middleware,
1358
+ context,
1293
1359
  })
1294
1360
 
1295
1361
  // Consume the stream to run the agentic loop