@tanstack/ai 0.6.2 → 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.
- package/dist/esm/activities/chat/index.d.ts +20 -0
- package/dist/esm/activities/chat/index.js +248 -213
- package/dist/esm/activities/chat/index.js.map +1 -1
- package/dist/esm/activities/chat/middleware/compose.d.ts +66 -0
- package/dist/esm/activities/chat/middleware/compose.js +327 -0
- package/dist/esm/activities/chat/middleware/compose.js.map +1 -0
- package/dist/esm/activities/chat/middleware/index.d.ts +2 -0
- package/dist/esm/activities/chat/middleware/tool-cache-middleware.d.ts +89 -0
- package/dist/esm/activities/chat/middleware/tool-cache-middleware.js +76 -0
- package/dist/esm/activities/chat/middleware/tool-cache-middleware.js.map +1 -0
- package/dist/esm/activities/chat/middleware/types.d.ts +307 -0
- package/dist/esm/activities/chat/tools/tool-calls.d.ts +16 -1
- package/dist/esm/activities/chat/tools/tool-calls.js +148 -64
- package/dist/esm/activities/chat/tools/tool-calls.js.map +1 -1
- package/dist/esm/activities/generateImage/index.d.ts +19 -6
- package/dist/esm/activities/generateImage/index.js +12 -3
- package/dist/esm/activities/generateImage/index.js.map +1 -1
- package/dist/esm/activities/generateSpeech/index.d.ts +19 -6
- package/dist/esm/activities/generateSpeech/index.js +12 -3
- package/dist/esm/activities/generateSpeech/index.js.map +1 -1
- package/dist/esm/activities/generateTranscription/index.d.ts +30 -6
- package/dist/esm/activities/generateTranscription/index.js +14 -3
- package/dist/esm/activities/generateTranscription/index.js.map +1 -1
- package/dist/esm/activities/generateVideo/index.d.ts +45 -7
- package/dist/esm/activities/generateVideo/index.js +91 -2
- package/dist/esm/activities/generateVideo/index.js.map +1 -1
- package/dist/esm/activities/stream-generation-result.d.ts +14 -0
- package/dist/esm/activities/stream-generation-result.js +40 -0
- package/dist/esm/activities/stream-generation-result.js.map +1 -0
- package/dist/esm/activities/summarize/index.js +3 -18
- package/dist/esm/activities/summarize/index.js.map +1 -1
- package/dist/esm/index.d.ts +3 -1
- package/dist/esm/index.js +2 -2
- package/dist/esm/middlewares/content-guard.d.ts +77 -0
- package/dist/esm/middlewares/content-guard.js +155 -0
- package/dist/esm/middlewares/content-guard.js.map +1 -0
- package/dist/esm/middlewares/index.d.ts +2 -0
- package/dist/esm/middlewares/index.js +7 -0
- package/dist/esm/middlewares/index.js.map +1 -0
- package/dist/esm/middlewares/tool-cache.d.ts +1 -0
- package/dist/esm/realtime/index.d.ts +30 -0
- package/dist/esm/realtime/index.js +8 -0
- package/dist/esm/realtime/index.js.map +1 -0
- package/dist/esm/realtime/types.d.ts +234 -0
- package/package.json +6 -6
- package/src/activities/chat/index.ts +322 -256
- package/src/activities/chat/middleware/compose.ts +392 -0
- package/src/activities/chat/middleware/index.ts +17 -0
- package/src/activities/chat/middleware/tool-cache-middleware.ts +189 -0
- package/src/activities/chat/middleware/types.ts +419 -0
- package/src/activities/chat/tools/tool-calls.ts +225 -87
- package/src/activities/generateImage/index.ts +50 -8
- package/src/activities/generateSpeech/index.ts +42 -8
- package/src/activities/generateTranscription/index.ts +60 -10
- package/src/activities/generateVideo/index.ts +174 -7
- package/src/activities/stream-generation-result.ts +62 -0
- package/src/activities/summarize/index.ts +4 -23
- package/src/index.ts +41 -2
- package/src/middlewares/content-guard.ts +285 -0
- package/src/middlewares/index.ts +13 -0
- package/src/middlewares/tool-cache.ts +6 -0
- package/src/realtime/index.ts +38 -0
- package/src/realtime/types.ts +294 -0
- package/dist/esm/event-client.d.ts +0 -394
- package/dist/esm/event-client.js +0 -13
- package/dist/esm/event-client.js.map +0 -1
- package/src/event-client.ts +0 -497
|
@@ -0,0 +1,392 @@
|
|
|
1
|
+
import { aiEventClient } from '@tanstack/ai-event-client'
|
|
2
|
+
import type { StreamChunk } from '../../../types'
|
|
3
|
+
import type {
|
|
4
|
+
AbortInfo,
|
|
5
|
+
AfterToolCallInfo,
|
|
6
|
+
BeforeToolCallDecision,
|
|
7
|
+
ChatMiddleware,
|
|
8
|
+
ChatMiddlewareConfig,
|
|
9
|
+
ChatMiddlewareContext,
|
|
10
|
+
ErrorInfo,
|
|
11
|
+
FinishInfo,
|
|
12
|
+
IterationInfo,
|
|
13
|
+
ToolCallHookContext,
|
|
14
|
+
ToolPhaseCompleteInfo,
|
|
15
|
+
UsageInfo,
|
|
16
|
+
} from './types'
|
|
17
|
+
|
|
18
|
+
/** Check if a middleware should be skipped for instrumentation events. */
|
|
19
|
+
function shouldSkipInstrumentation(mw: ChatMiddleware): boolean {
|
|
20
|
+
return mw.name === 'devtools'
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
/** Build the base context for middleware instrumentation events. */
|
|
24
|
+
function instrumentCtx(ctx: ChatMiddlewareContext) {
|
|
25
|
+
return {
|
|
26
|
+
requestId: ctx.requestId,
|
|
27
|
+
streamId: ctx.streamId,
|
|
28
|
+
clientId: ctx.conversationId,
|
|
29
|
+
timestamp: Date.now(),
|
|
30
|
+
}
|
|
31
|
+
}
|
|
32
|
+
|
|
33
|
+
/**
|
|
34
|
+
* Internal middleware runner that manages composed execution of middleware hooks.
|
|
35
|
+
* Created once per chat() invocation.
|
|
36
|
+
*/
|
|
37
|
+
export class MiddlewareRunner {
|
|
38
|
+
private readonly middlewares: ReadonlyArray<ChatMiddleware>
|
|
39
|
+
|
|
40
|
+
constructor(middlewares: ReadonlyArray<ChatMiddleware>) {
|
|
41
|
+
this.middlewares = middlewares
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
get hasMiddleware(): boolean {
|
|
45
|
+
return this.middlewares.length > 0
|
|
46
|
+
}
|
|
47
|
+
|
|
48
|
+
/**
|
|
49
|
+
* Pipe config through all middleware onConfig hooks in order.
|
|
50
|
+
* Each middleware receives the merged config from previous middleware.
|
|
51
|
+
* Partial returns are shallow-merged with the current config.
|
|
52
|
+
*/
|
|
53
|
+
async runOnConfig(
|
|
54
|
+
ctx: ChatMiddlewareContext,
|
|
55
|
+
config: ChatMiddlewareConfig,
|
|
56
|
+
): Promise<ChatMiddlewareConfig> {
|
|
57
|
+
let current = config
|
|
58
|
+
for (const mw of this.middlewares) {
|
|
59
|
+
if (mw.onConfig) {
|
|
60
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
61
|
+
const start = Date.now()
|
|
62
|
+
const result = await mw.onConfig(ctx, current)
|
|
63
|
+
const hasTransform = result !== undefined && result !== null
|
|
64
|
+
if (hasTransform) {
|
|
65
|
+
current = { ...current, ...result }
|
|
66
|
+
}
|
|
67
|
+
if (!skip) {
|
|
68
|
+
const base = instrumentCtx(ctx)
|
|
69
|
+
aiEventClient.emit('middleware:hook:executed', {
|
|
70
|
+
...base,
|
|
71
|
+
middlewareName: mw.name || 'unnamed',
|
|
72
|
+
hookName: 'onConfig',
|
|
73
|
+
iteration: ctx.iteration,
|
|
74
|
+
duration: Date.now() - start,
|
|
75
|
+
hasTransform,
|
|
76
|
+
})
|
|
77
|
+
if (hasTransform) {
|
|
78
|
+
aiEventClient.emit('middleware:config:transformed', {
|
|
79
|
+
...base,
|
|
80
|
+
middlewareName: mw.name || 'unnamed',
|
|
81
|
+
iteration: ctx.iteration,
|
|
82
|
+
changes: result as Record<string, unknown>,
|
|
83
|
+
})
|
|
84
|
+
}
|
|
85
|
+
}
|
|
86
|
+
}
|
|
87
|
+
}
|
|
88
|
+
return current
|
|
89
|
+
}
|
|
90
|
+
|
|
91
|
+
/**
|
|
92
|
+
* Call onStart on all middleware in order.
|
|
93
|
+
*/
|
|
94
|
+
async runOnStart(ctx: ChatMiddlewareContext): Promise<void> {
|
|
95
|
+
for (const mw of this.middlewares) {
|
|
96
|
+
if (mw.onStart) {
|
|
97
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
98
|
+
const start = Date.now()
|
|
99
|
+
await mw.onStart(ctx)
|
|
100
|
+
if (!skip) {
|
|
101
|
+
aiEventClient.emit('middleware:hook:executed', {
|
|
102
|
+
...instrumentCtx(ctx),
|
|
103
|
+
middlewareName: mw.name || 'unnamed',
|
|
104
|
+
hookName: 'onStart',
|
|
105
|
+
iteration: ctx.iteration,
|
|
106
|
+
duration: Date.now() - start,
|
|
107
|
+
hasTransform: false,
|
|
108
|
+
})
|
|
109
|
+
}
|
|
110
|
+
}
|
|
111
|
+
}
|
|
112
|
+
}
|
|
113
|
+
|
|
114
|
+
/**
|
|
115
|
+
* Pipe a single chunk through all middleware onChunk hooks in order.
|
|
116
|
+
* Returns the resulting chunks (0..N) to yield to the consumer.
|
|
117
|
+
*
|
|
118
|
+
* - void: pass through unchanged
|
|
119
|
+
* - chunk: replace with this chunk
|
|
120
|
+
* - chunk[]: expand to multiple chunks
|
|
121
|
+
* - null: drop the chunk entirely
|
|
122
|
+
*/
|
|
123
|
+
async runOnChunk(
|
|
124
|
+
ctx: ChatMiddlewareContext,
|
|
125
|
+
chunk: StreamChunk,
|
|
126
|
+
): Promise<Array<StreamChunk>> {
|
|
127
|
+
let chunks: Array<StreamChunk> = [chunk]
|
|
128
|
+
|
|
129
|
+
for (const mw of this.middlewares) {
|
|
130
|
+
if (!mw.onChunk) continue
|
|
131
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
132
|
+
|
|
133
|
+
const nextChunks: Array<StreamChunk> = []
|
|
134
|
+
for (const c of chunks) {
|
|
135
|
+
const result = await mw.onChunk(ctx, c)
|
|
136
|
+
if (result === null) {
|
|
137
|
+
// Drop this chunk
|
|
138
|
+
if (!skip) {
|
|
139
|
+
aiEventClient.emit('middleware:chunk:transformed', {
|
|
140
|
+
...instrumentCtx(ctx),
|
|
141
|
+
middlewareName: mw.name || 'unnamed',
|
|
142
|
+
originalChunkType: c.type,
|
|
143
|
+
resultCount: 0,
|
|
144
|
+
wasDropped: true,
|
|
145
|
+
})
|
|
146
|
+
}
|
|
147
|
+
continue
|
|
148
|
+
} else if (result === undefined) {
|
|
149
|
+
// Pass through — no instrumentation for pass-throughs
|
|
150
|
+
nextChunks.push(c)
|
|
151
|
+
} else if (Array.isArray(result)) {
|
|
152
|
+
// Expand
|
|
153
|
+
nextChunks.push(...result)
|
|
154
|
+
if (!skip) {
|
|
155
|
+
aiEventClient.emit('middleware:chunk:transformed', {
|
|
156
|
+
...instrumentCtx(ctx),
|
|
157
|
+
middlewareName: mw.name || 'unnamed',
|
|
158
|
+
originalChunkType: c.type,
|
|
159
|
+
resultCount: result.length,
|
|
160
|
+
wasDropped: false,
|
|
161
|
+
})
|
|
162
|
+
}
|
|
163
|
+
} else {
|
|
164
|
+
// Replace
|
|
165
|
+
nextChunks.push(result)
|
|
166
|
+
if (!skip) {
|
|
167
|
+
aiEventClient.emit('middleware:chunk:transformed', {
|
|
168
|
+
...instrumentCtx(ctx),
|
|
169
|
+
middlewareName: mw.name || 'unnamed',
|
|
170
|
+
originalChunkType: c.type,
|
|
171
|
+
resultCount: 1,
|
|
172
|
+
wasDropped: false,
|
|
173
|
+
})
|
|
174
|
+
}
|
|
175
|
+
}
|
|
176
|
+
}
|
|
177
|
+
chunks = nextChunks
|
|
178
|
+
}
|
|
179
|
+
|
|
180
|
+
return chunks
|
|
181
|
+
}
|
|
182
|
+
|
|
183
|
+
/**
|
|
184
|
+
* Run onBeforeToolCall through middleware in order.
|
|
185
|
+
* Returns the first non-void decision, or undefined to continue normally.
|
|
186
|
+
*/
|
|
187
|
+
async runOnBeforeToolCall(
|
|
188
|
+
ctx: ChatMiddlewareContext,
|
|
189
|
+
hookCtx: ToolCallHookContext,
|
|
190
|
+
): Promise<BeforeToolCallDecision> {
|
|
191
|
+
for (const mw of this.middlewares) {
|
|
192
|
+
if (mw.onBeforeToolCall) {
|
|
193
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
194
|
+
const start = Date.now()
|
|
195
|
+
const decision = await mw.onBeforeToolCall(ctx, hookCtx)
|
|
196
|
+
const hasTransform = decision !== undefined && decision !== null
|
|
197
|
+
if (!skip) {
|
|
198
|
+
aiEventClient.emit('middleware:hook:executed', {
|
|
199
|
+
...instrumentCtx(ctx),
|
|
200
|
+
middlewareName: mw.name || 'unnamed',
|
|
201
|
+
hookName: 'onBeforeToolCall',
|
|
202
|
+
iteration: ctx.iteration,
|
|
203
|
+
duration: Date.now() - start,
|
|
204
|
+
hasTransform,
|
|
205
|
+
})
|
|
206
|
+
}
|
|
207
|
+
if (hasTransform) {
|
|
208
|
+
return decision
|
|
209
|
+
}
|
|
210
|
+
}
|
|
211
|
+
}
|
|
212
|
+
return undefined
|
|
213
|
+
}
|
|
214
|
+
|
|
215
|
+
/**
|
|
216
|
+
* Run onAfterToolCall on all middleware in order.
|
|
217
|
+
*/
|
|
218
|
+
async runOnAfterToolCall(
|
|
219
|
+
ctx: ChatMiddlewareContext,
|
|
220
|
+
info: AfterToolCallInfo,
|
|
221
|
+
): Promise<void> {
|
|
222
|
+
for (const mw of this.middlewares) {
|
|
223
|
+
if (mw.onAfterToolCall) {
|
|
224
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
225
|
+
const start = Date.now()
|
|
226
|
+
await mw.onAfterToolCall(ctx, info)
|
|
227
|
+
if (!skip) {
|
|
228
|
+
aiEventClient.emit('middleware:hook:executed', {
|
|
229
|
+
...instrumentCtx(ctx),
|
|
230
|
+
middlewareName: mw.name || 'unnamed',
|
|
231
|
+
hookName: 'onAfterToolCall',
|
|
232
|
+
iteration: ctx.iteration,
|
|
233
|
+
duration: Date.now() - start,
|
|
234
|
+
hasTransform: false,
|
|
235
|
+
})
|
|
236
|
+
}
|
|
237
|
+
}
|
|
238
|
+
}
|
|
239
|
+
}
|
|
240
|
+
|
|
241
|
+
/**
|
|
242
|
+
* Run onUsage on all middleware in order.
|
|
243
|
+
*/
|
|
244
|
+
async runOnUsage(
|
|
245
|
+
ctx: ChatMiddlewareContext,
|
|
246
|
+
usage: UsageInfo,
|
|
247
|
+
): Promise<void> {
|
|
248
|
+
for (const mw of this.middlewares) {
|
|
249
|
+
if (mw.onUsage) {
|
|
250
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
251
|
+
const start = Date.now()
|
|
252
|
+
await mw.onUsage(ctx, usage)
|
|
253
|
+
if (!skip) {
|
|
254
|
+
aiEventClient.emit('middleware:hook:executed', {
|
|
255
|
+
...instrumentCtx(ctx),
|
|
256
|
+
middlewareName: mw.name || 'unnamed',
|
|
257
|
+
hookName: 'onUsage',
|
|
258
|
+
iteration: ctx.iteration,
|
|
259
|
+
duration: Date.now() - start,
|
|
260
|
+
hasTransform: false,
|
|
261
|
+
})
|
|
262
|
+
}
|
|
263
|
+
}
|
|
264
|
+
}
|
|
265
|
+
}
|
|
266
|
+
|
|
267
|
+
/**
|
|
268
|
+
* Run onFinish on all middleware in order.
|
|
269
|
+
*/
|
|
270
|
+
async runOnFinish(
|
|
271
|
+
ctx: ChatMiddlewareContext,
|
|
272
|
+
info: FinishInfo,
|
|
273
|
+
): Promise<void> {
|
|
274
|
+
for (const mw of this.middlewares) {
|
|
275
|
+
if (mw.onFinish) {
|
|
276
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
277
|
+
const start = Date.now()
|
|
278
|
+
await mw.onFinish(ctx, info)
|
|
279
|
+
if (!skip) {
|
|
280
|
+
aiEventClient.emit('middleware:hook:executed', {
|
|
281
|
+
...instrumentCtx(ctx),
|
|
282
|
+
middlewareName: mw.name || 'unnamed',
|
|
283
|
+
hookName: 'onFinish',
|
|
284
|
+
iteration: ctx.iteration,
|
|
285
|
+
duration: Date.now() - start,
|
|
286
|
+
hasTransform: false,
|
|
287
|
+
})
|
|
288
|
+
}
|
|
289
|
+
}
|
|
290
|
+
}
|
|
291
|
+
}
|
|
292
|
+
|
|
293
|
+
/**
|
|
294
|
+
* Run onAbort on all middleware in order.
|
|
295
|
+
*/
|
|
296
|
+
async runOnAbort(ctx: ChatMiddlewareContext, info: AbortInfo): Promise<void> {
|
|
297
|
+
for (const mw of this.middlewares) {
|
|
298
|
+
if (mw.onAbort) {
|
|
299
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
300
|
+
const start = Date.now()
|
|
301
|
+
await mw.onAbort(ctx, info)
|
|
302
|
+
if (!skip) {
|
|
303
|
+
aiEventClient.emit('middleware:hook:executed', {
|
|
304
|
+
...instrumentCtx(ctx),
|
|
305
|
+
middlewareName: mw.name || 'unnamed',
|
|
306
|
+
hookName: 'onAbort',
|
|
307
|
+
iteration: ctx.iteration,
|
|
308
|
+
duration: Date.now() - start,
|
|
309
|
+
hasTransform: false,
|
|
310
|
+
})
|
|
311
|
+
}
|
|
312
|
+
}
|
|
313
|
+
}
|
|
314
|
+
}
|
|
315
|
+
|
|
316
|
+
/**
|
|
317
|
+
* Run onError on all middleware in order.
|
|
318
|
+
*/
|
|
319
|
+
async runOnError(ctx: ChatMiddlewareContext, info: ErrorInfo): Promise<void> {
|
|
320
|
+
for (const mw of this.middlewares) {
|
|
321
|
+
if (mw.onError) {
|
|
322
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
323
|
+
const start = Date.now()
|
|
324
|
+
await mw.onError(ctx, info)
|
|
325
|
+
if (!skip) {
|
|
326
|
+
aiEventClient.emit('middleware:hook:executed', {
|
|
327
|
+
...instrumentCtx(ctx),
|
|
328
|
+
middlewareName: mw.name || 'unnamed',
|
|
329
|
+
hookName: 'onError',
|
|
330
|
+
iteration: ctx.iteration,
|
|
331
|
+
duration: Date.now() - start,
|
|
332
|
+
hasTransform: false,
|
|
333
|
+
})
|
|
334
|
+
}
|
|
335
|
+
}
|
|
336
|
+
}
|
|
337
|
+
}
|
|
338
|
+
|
|
339
|
+
/**
|
|
340
|
+
* Run onIteration on all middleware in order.
|
|
341
|
+
* Called at the start of each agent loop iteration.
|
|
342
|
+
*/
|
|
343
|
+
async runOnIteration(
|
|
344
|
+
ctx: ChatMiddlewareContext,
|
|
345
|
+
info: IterationInfo,
|
|
346
|
+
): Promise<void> {
|
|
347
|
+
for (const mw of this.middlewares) {
|
|
348
|
+
if (mw.onIteration) {
|
|
349
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
350
|
+
const start = Date.now()
|
|
351
|
+
await mw.onIteration(ctx, info)
|
|
352
|
+
if (!skip) {
|
|
353
|
+
aiEventClient.emit('middleware:hook:executed', {
|
|
354
|
+
...instrumentCtx(ctx),
|
|
355
|
+
middlewareName: mw.name || 'unnamed',
|
|
356
|
+
hookName: 'onIteration',
|
|
357
|
+
iteration: ctx.iteration,
|
|
358
|
+
duration: Date.now() - start,
|
|
359
|
+
hasTransform: false,
|
|
360
|
+
})
|
|
361
|
+
}
|
|
362
|
+
}
|
|
363
|
+
}
|
|
364
|
+
}
|
|
365
|
+
|
|
366
|
+
/**
|
|
367
|
+
* Run onToolPhaseComplete on all middleware in order.
|
|
368
|
+
* Called after all tool calls in an iteration have been processed.
|
|
369
|
+
*/
|
|
370
|
+
async runOnToolPhaseComplete(
|
|
371
|
+
ctx: ChatMiddlewareContext,
|
|
372
|
+
info: ToolPhaseCompleteInfo,
|
|
373
|
+
): Promise<void> {
|
|
374
|
+
for (const mw of this.middlewares) {
|
|
375
|
+
if (mw.onToolPhaseComplete) {
|
|
376
|
+
const skip = shouldSkipInstrumentation(mw)
|
|
377
|
+
const start = Date.now()
|
|
378
|
+
await mw.onToolPhaseComplete(ctx, info)
|
|
379
|
+
if (!skip) {
|
|
380
|
+
aiEventClient.emit('middleware:hook:executed', {
|
|
381
|
+
...instrumentCtx(ctx),
|
|
382
|
+
middlewareName: mw.name || 'unnamed',
|
|
383
|
+
hookName: 'onToolPhaseComplete',
|
|
384
|
+
iteration: ctx.iteration,
|
|
385
|
+
duration: Date.now() - start,
|
|
386
|
+
hasTransform: false,
|
|
387
|
+
})
|
|
388
|
+
}
|
|
389
|
+
}
|
|
390
|
+
}
|
|
391
|
+
}
|
|
392
|
+
}
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
export type {
|
|
2
|
+
ChatMiddleware,
|
|
3
|
+
ChatMiddlewareContext,
|
|
4
|
+
ChatMiddlewarePhase,
|
|
5
|
+
ChatMiddlewareConfig,
|
|
6
|
+
ToolCallHookContext,
|
|
7
|
+
BeforeToolCallDecision,
|
|
8
|
+
AfterToolCallInfo,
|
|
9
|
+
IterationInfo,
|
|
10
|
+
ToolPhaseCompleteInfo,
|
|
11
|
+
UsageInfo,
|
|
12
|
+
FinishInfo,
|
|
13
|
+
AbortInfo,
|
|
14
|
+
ErrorInfo,
|
|
15
|
+
} from './types'
|
|
16
|
+
|
|
17
|
+
export { MiddlewareRunner } from './compose'
|
|
@@ -0,0 +1,189 @@
|
|
|
1
|
+
import type { ChatMiddleware } from './types'
|
|
2
|
+
|
|
3
|
+
/**
|
|
4
|
+
* A cache entry stored by the tool cache middleware.
|
|
5
|
+
*/
|
|
6
|
+
export interface ToolCacheEntry {
|
|
7
|
+
result: unknown
|
|
8
|
+
timestamp: number
|
|
9
|
+
}
|
|
10
|
+
|
|
11
|
+
/**
|
|
12
|
+
* Custom storage backend for the tool cache middleware.
|
|
13
|
+
*
|
|
14
|
+
* When provided, the middleware delegates all cache operations to this storage
|
|
15
|
+
* instead of using the built-in in-memory Map. This enables external storage
|
|
16
|
+
* backends like Redis, localStorage, databases, etc.
|
|
17
|
+
*
|
|
18
|
+
* All methods may return a Promise for async storage backends.
|
|
19
|
+
*/
|
|
20
|
+
export interface ToolCacheStorage {
|
|
21
|
+
getItem: (
|
|
22
|
+
key: string,
|
|
23
|
+
) => ToolCacheEntry | undefined | Promise<ToolCacheEntry | undefined>
|
|
24
|
+
setItem: (key: string, value: ToolCacheEntry) => void | Promise<void>
|
|
25
|
+
deleteItem: (key: string) => void | Promise<void>
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
/**
|
|
29
|
+
* Options for the tool cache middleware.
|
|
30
|
+
*/
|
|
31
|
+
export interface ToolCacheMiddlewareOptions {
|
|
32
|
+
/**
|
|
33
|
+
* Maximum number of entries in the cache.
|
|
34
|
+
* When exceeded, the oldest entry is evicted (LRU).
|
|
35
|
+
*
|
|
36
|
+
* Only applies to the default in-memory storage.
|
|
37
|
+
* When a custom `storage` is provided, capacity management is the storage's responsibility.
|
|
38
|
+
*
|
|
39
|
+
* @default 100
|
|
40
|
+
*/
|
|
41
|
+
maxSize?: number
|
|
42
|
+
|
|
43
|
+
/**
|
|
44
|
+
* Time-to-live in milliseconds. Entries older than this are not served from cache.
|
|
45
|
+
* @default Infinity (no expiry)
|
|
46
|
+
*/
|
|
47
|
+
ttl?: number
|
|
48
|
+
|
|
49
|
+
/**
|
|
50
|
+
* Tool names to cache. If not provided, all tools are cached.
|
|
51
|
+
*/
|
|
52
|
+
toolNames?: Array<string>
|
|
53
|
+
|
|
54
|
+
/**
|
|
55
|
+
* Custom function to generate a cache key from tool name and args.
|
|
56
|
+
* Defaults to `JSON.stringify([toolName, args])`.
|
|
57
|
+
*/
|
|
58
|
+
keyFn?: (toolName: string, args: unknown) => string
|
|
59
|
+
|
|
60
|
+
/**
|
|
61
|
+
* Custom storage backend. When provided, the middleware uses this instead of
|
|
62
|
+
* the built-in in-memory Map. The storage is responsible for its own capacity
|
|
63
|
+
* management — the `maxSize` option is ignored.
|
|
64
|
+
*
|
|
65
|
+
* @example
|
|
66
|
+
* ```ts
|
|
67
|
+
* toolCacheMiddleware({
|
|
68
|
+
* storage: {
|
|
69
|
+
* getItem: (key) => redisClient.get(key).then(v => v ? JSON.parse(v) : undefined),
|
|
70
|
+
* setItem: (key, value) => redisClient.set(key, JSON.stringify(value)),
|
|
71
|
+
* deleteItem: (key) => redisClient.del(key),
|
|
72
|
+
* },
|
|
73
|
+
* })
|
|
74
|
+
* ```
|
|
75
|
+
*/
|
|
76
|
+
storage?: ToolCacheStorage
|
|
77
|
+
}
|
|
78
|
+
|
|
79
|
+
function defaultKeyFn(toolName: string, args: unknown): string {
|
|
80
|
+
return JSON.stringify([toolName, args])
|
|
81
|
+
}
|
|
82
|
+
|
|
83
|
+
function createDefaultStorage(maxSize: number): ToolCacheStorage {
|
|
84
|
+
const cache = new Map<string, ToolCacheEntry>()
|
|
85
|
+
|
|
86
|
+
return {
|
|
87
|
+
getItem: (key) => {
|
|
88
|
+
const entry = cache.get(key)
|
|
89
|
+
if (entry !== undefined) {
|
|
90
|
+
// Refresh recency: delete and re-insert so this key becomes newest
|
|
91
|
+
cache.delete(key)
|
|
92
|
+
cache.set(key, entry)
|
|
93
|
+
}
|
|
94
|
+
return entry
|
|
95
|
+
},
|
|
96
|
+
setItem: (key, value) => {
|
|
97
|
+
// Delete first so re-inserts also refresh recency
|
|
98
|
+
if (cache.has(key)) {
|
|
99
|
+
cache.delete(key)
|
|
100
|
+
} else if (cache.size >= maxSize) {
|
|
101
|
+
// LRU eviction: Map iteration order is insertion order — first key is least recently used
|
|
102
|
+
const firstKey = cache.keys().next().value
|
|
103
|
+
if (firstKey !== undefined) {
|
|
104
|
+
cache.delete(firstKey)
|
|
105
|
+
}
|
|
106
|
+
}
|
|
107
|
+
cache.set(key, value)
|
|
108
|
+
},
|
|
109
|
+
deleteItem: (key) => {
|
|
110
|
+
cache.delete(key)
|
|
111
|
+
},
|
|
112
|
+
}
|
|
113
|
+
}
|
|
114
|
+
|
|
115
|
+
/**
|
|
116
|
+
* Creates a middleware that caches tool call results based on tool name + arguments.
|
|
117
|
+
*
|
|
118
|
+
* When a tool is called with the same name and arguments as a previous call,
|
|
119
|
+
* the cached result is returned immediately without executing the tool.
|
|
120
|
+
*
|
|
121
|
+
* @example
|
|
122
|
+
* ```ts
|
|
123
|
+
* import { chat, toolCacheMiddleware } from '@tanstack/ai'
|
|
124
|
+
*
|
|
125
|
+
* const stream = chat({
|
|
126
|
+
* adapter,
|
|
127
|
+
* messages,
|
|
128
|
+
* tools: [weatherTool, stockTool],
|
|
129
|
+
* middleware: [
|
|
130
|
+
* toolCacheMiddleware({ ttl: 60_000, toolNames: ['getWeather'] }),
|
|
131
|
+
* ],
|
|
132
|
+
* })
|
|
133
|
+
* ```
|
|
134
|
+
*/
|
|
135
|
+
export function toolCacheMiddleware(
|
|
136
|
+
options: ToolCacheMiddlewareOptions = {},
|
|
137
|
+
): ChatMiddleware {
|
|
138
|
+
const {
|
|
139
|
+
maxSize = 100,
|
|
140
|
+
ttl = Infinity,
|
|
141
|
+
toolNames,
|
|
142
|
+
keyFn = defaultKeyFn,
|
|
143
|
+
storage = createDefaultStorage(maxSize),
|
|
144
|
+
} = options
|
|
145
|
+
|
|
146
|
+
return {
|
|
147
|
+
name: 'tool-cache-middleware',
|
|
148
|
+
|
|
149
|
+
onBeforeToolCall: async (_ctx, hookCtx) => {
|
|
150
|
+
if (toolNames && !toolNames.includes(hookCtx.toolName)) {
|
|
151
|
+
return undefined
|
|
152
|
+
}
|
|
153
|
+
|
|
154
|
+
const key = keyFn(hookCtx.toolName, hookCtx.args)
|
|
155
|
+
const entry = await storage.getItem(key)
|
|
156
|
+
|
|
157
|
+
if (entry) {
|
|
158
|
+
const age = Date.now() - entry.timestamp
|
|
159
|
+
if (age < ttl) {
|
|
160
|
+
return { type: 'skip', result: entry.result }
|
|
161
|
+
}
|
|
162
|
+
// Expired — remove
|
|
163
|
+
await storage.deleteItem(key)
|
|
164
|
+
}
|
|
165
|
+
|
|
166
|
+
return undefined
|
|
167
|
+
},
|
|
168
|
+
|
|
169
|
+
onAfterToolCall: async (_ctx, info) => {
|
|
170
|
+
if (!info.ok) return
|
|
171
|
+
if (toolNames && !toolNames.includes(info.toolName)) return
|
|
172
|
+
|
|
173
|
+
// Re-derive the key from the raw arguments to match what onBeforeToolCall produces
|
|
174
|
+
let parsedArgs: unknown
|
|
175
|
+
try {
|
|
176
|
+
parsedArgs = JSON.parse(info.toolCall.function.arguments.trim() || '{}')
|
|
177
|
+
} catch {
|
|
178
|
+
return
|
|
179
|
+
}
|
|
180
|
+
|
|
181
|
+
const key = keyFn(info.toolName, parsedArgs)
|
|
182
|
+
|
|
183
|
+
await storage.setItem(key, {
|
|
184
|
+
result: info.result,
|
|
185
|
+
timestamp: Date.now(),
|
|
186
|
+
})
|
|
187
|
+
},
|
|
188
|
+
}
|
|
189
|
+
}
|