@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.
- 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.js +1 -1
- package/dist/esm/activities/generateImage/index.js.map +1 -1
- package/dist/esm/activities/generateSpeech/index.js +1 -1
- package/dist/esm/activities/generateSpeech/index.js.map +1 -1
- package/dist/esm/activities/generateTranscription/index.js +1 -1
- package/dist/esm/activities/generateTranscription/index.js.map +1 -1
- package/dist/esm/activities/generateVideo/index.js +1 -1
- package/dist/esm/activities/generateVideo/index.js.map +1 -1
- package/dist/esm/activities/summarize/index.js +1 -1
- 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 +1 -1
- package/src/activities/generateSpeech/index.ts +1 -1
- package/src/activities/generateTranscription/index.ts +1 -1
- package/src/activities/generateVideo/index.ts +1 -1
- package/src/activities/summarize/index.ts +1 -1
- 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
|
@@ -1,9 +1,10 @@
|
|
|
1
|
-
import {
|
|
1
|
+
import { devtoolsMiddleware } from "@tanstack/ai-event-client";
|
|
2
2
|
import { streamToText } from "../../stream-to-response.js";
|
|
3
|
-
import { ToolCallManager, executeToolCalls } from "./tools/tool-calls.js";
|
|
3
|
+
import { ToolCallManager, MiddlewareAbortError, executeToolCalls } from "./tools/tool-calls.js";
|
|
4
4
|
import { convertSchemaToJsonSchema, isStandardSchema, parseWithStandardSchema } from "./tools/schema-converter.js";
|
|
5
5
|
import { maxIterations } from "./agent-loop-strategies.js";
|
|
6
6
|
import { convertMessagesToModelMessages } from "./messages.js";
|
|
7
|
+
import { MiddlewareRunner } from "./middleware/compose.js";
|
|
7
8
|
const kind = "text";
|
|
8
9
|
function createChatOptions(options) {
|
|
9
10
|
return options;
|
|
@@ -17,10 +18,11 @@ class TextEngine {
|
|
|
17
18
|
this.currentMessageId = null;
|
|
18
19
|
this.accumulatedContent = "";
|
|
19
20
|
this.finishedEvent = null;
|
|
20
|
-
this.shouldEmitStreamEnd = true;
|
|
21
21
|
this.earlyTermination = false;
|
|
22
22
|
this.toolPhase = "continue";
|
|
23
23
|
this.cyclePhase = "processText";
|
|
24
|
+
this.deferredPromises = [];
|
|
25
|
+
this.terminalHookCalled = false;
|
|
24
26
|
this.adapter = config.adapter;
|
|
25
27
|
this.params = config.params;
|
|
26
28
|
this.systemPrompts = config.params.systemPrompts || [];
|
|
@@ -40,6 +42,45 @@ class TextEngine {
|
|
|
40
42
|
this.streamId = this.createId("stream");
|
|
41
43
|
this.effectiveRequest = config.params.abortController ? { signal: config.params.abortController.signal } : void 0;
|
|
42
44
|
this.effectiveSignal = config.params.abortController?.signal;
|
|
45
|
+
const allMiddleware = [devtoolsMiddleware(), ...config.middleware || []];
|
|
46
|
+
this.middlewareRunner = new MiddlewareRunner(allMiddleware);
|
|
47
|
+
this.middlewareAbortController = new AbortController();
|
|
48
|
+
this.middlewareCtx = {
|
|
49
|
+
requestId: this.requestId,
|
|
50
|
+
streamId: this.streamId,
|
|
51
|
+
conversationId: config.params.conversationId,
|
|
52
|
+
phase: "init",
|
|
53
|
+
iteration: 0,
|
|
54
|
+
chunkIndex: 0,
|
|
55
|
+
signal: this.effectiveSignal,
|
|
56
|
+
abort: (reason) => {
|
|
57
|
+
this.abortReason = reason;
|
|
58
|
+
this.middlewareAbortController?.abort(reason);
|
|
59
|
+
},
|
|
60
|
+
context: config.context,
|
|
61
|
+
defer: (promise) => {
|
|
62
|
+
this.deferredPromises.push(promise);
|
|
63
|
+
},
|
|
64
|
+
// Provider / adapter info
|
|
65
|
+
provider: config.adapter.name,
|
|
66
|
+
model: config.params.model,
|
|
67
|
+
source: "server",
|
|
68
|
+
streaming: true,
|
|
69
|
+
// Config-derived (updated in beforeRun and applyMiddlewareConfig)
|
|
70
|
+
systemPrompts: this.systemPrompts,
|
|
71
|
+
toolNames: void 0,
|
|
72
|
+
options: void 0,
|
|
73
|
+
modelOptions: config.params.modelOptions,
|
|
74
|
+
// Computed
|
|
75
|
+
messageCount: this.initialMessageCount,
|
|
76
|
+
hasTools: this.tools.length > 0,
|
|
77
|
+
// Mutable per-iteration
|
|
78
|
+
currentMessageId: null,
|
|
79
|
+
accumulatedContent: "",
|
|
80
|
+
// References
|
|
81
|
+
messages: this.messages,
|
|
82
|
+
createId: (prefix) => this.createId(prefix)
|
|
83
|
+
};
|
|
43
84
|
}
|
|
44
85
|
/** Get the accumulated content after the chat loop completes */
|
|
45
86
|
getAccumulatedContent() {
|
|
@@ -52,24 +93,77 @@ class TextEngine {
|
|
|
52
93
|
async *run() {
|
|
53
94
|
this.beforeRun();
|
|
54
95
|
try {
|
|
96
|
+
this.middlewareCtx.phase = "init";
|
|
97
|
+
const initialConfig = this.buildMiddlewareConfig();
|
|
98
|
+
const transformedConfig = await this.middlewareRunner.runOnConfig(
|
|
99
|
+
this.middlewareCtx,
|
|
100
|
+
initialConfig
|
|
101
|
+
);
|
|
102
|
+
this.applyMiddlewareConfig(transformedConfig);
|
|
103
|
+
await this.middlewareRunner.runOnStart(this.middlewareCtx);
|
|
55
104
|
const pendingPhase = yield* this.checkForPendingToolCalls();
|
|
56
105
|
if (pendingPhase === "wait") {
|
|
57
106
|
return;
|
|
58
107
|
}
|
|
59
108
|
do {
|
|
60
|
-
if (this.earlyTermination || this.
|
|
109
|
+
if (this.earlyTermination || this.isCancelled()) {
|
|
61
110
|
return;
|
|
62
111
|
}
|
|
63
|
-
this.beginCycle();
|
|
112
|
+
await this.beginCycle();
|
|
64
113
|
if (this.cyclePhase === "processText") {
|
|
114
|
+
this.middlewareCtx.phase = "beforeModel";
|
|
115
|
+
this.middlewareCtx.iteration = this.iterationCount;
|
|
116
|
+
const iterConfig = this.buildMiddlewareConfig();
|
|
117
|
+
const transformedConfig2 = await this.middlewareRunner.runOnConfig(
|
|
118
|
+
this.middlewareCtx,
|
|
119
|
+
iterConfig
|
|
120
|
+
);
|
|
121
|
+
this.applyMiddlewareConfig(transformedConfig2);
|
|
65
122
|
yield* this.streamModelResponse();
|
|
66
123
|
} else {
|
|
67
124
|
yield* this.processToolCalls();
|
|
68
125
|
}
|
|
69
126
|
this.endCycle();
|
|
70
127
|
} while (this.shouldContinue());
|
|
128
|
+
if (!this.terminalHookCalled && this.toolPhase !== "wait") {
|
|
129
|
+
this.terminalHookCalled = true;
|
|
130
|
+
await this.middlewareRunner.runOnFinish(this.middlewareCtx, {
|
|
131
|
+
finishReason: this.lastFinishReason,
|
|
132
|
+
duration: Date.now() - this.streamStartTime,
|
|
133
|
+
content: this.accumulatedContent,
|
|
134
|
+
usage: this.finishedEvent?.usage
|
|
135
|
+
});
|
|
136
|
+
}
|
|
137
|
+
} catch (error) {
|
|
138
|
+
if (!this.terminalHookCalled) {
|
|
139
|
+
this.terminalHookCalled = true;
|
|
140
|
+
if (error instanceof MiddlewareAbortError) {
|
|
141
|
+
this.abortReason = error.message;
|
|
142
|
+
await this.middlewareRunner.runOnAbort(this.middlewareCtx, {
|
|
143
|
+
reason: error.message,
|
|
144
|
+
duration: Date.now() - this.streamStartTime
|
|
145
|
+
});
|
|
146
|
+
} else {
|
|
147
|
+
await this.middlewareRunner.runOnError(this.middlewareCtx, {
|
|
148
|
+
error,
|
|
149
|
+
duration: Date.now() - this.streamStartTime
|
|
150
|
+
});
|
|
151
|
+
}
|
|
152
|
+
}
|
|
153
|
+
if (!(error instanceof MiddlewareAbortError)) {
|
|
154
|
+
throw error;
|
|
155
|
+
}
|
|
71
156
|
} finally {
|
|
72
|
-
this.
|
|
157
|
+
if (!this.terminalHookCalled && this.isCancelled()) {
|
|
158
|
+
this.terminalHookCalled = true;
|
|
159
|
+
await this.middlewareRunner.runOnAbort(this.middlewareCtx, {
|
|
160
|
+
reason: this.abortReason,
|
|
161
|
+
duration: Date.now() - this.streamStartTime
|
|
162
|
+
});
|
|
163
|
+
}
|
|
164
|
+
if (this.deferredPromises.length > 0) {
|
|
165
|
+
await Promise.allSettled(this.deferredPromises);
|
|
166
|
+
}
|
|
73
167
|
}
|
|
74
168
|
}
|
|
75
169
|
beforeRun() {
|
|
@@ -82,55 +176,12 @@ class TextEngine {
|
|
|
82
176
|
if (metadata !== void 0) options.metadata = metadata;
|
|
83
177
|
this.eventOptions = Object.keys(options).length > 0 ? options : void 0;
|
|
84
178
|
this.eventToolNames = tools?.map((t) => t.name);
|
|
85
|
-
|
|
86
|
-
|
|
87
|
-
timestamp: Date.now()
|
|
88
|
-
});
|
|
89
|
-
const messagesToEmit = this.params.conversationId ? this.messages.slice(-1).filter((m) => m.role === "user") : this.messages;
|
|
90
|
-
messagesToEmit.forEach((message, index) => {
|
|
91
|
-
const messageIndex = this.params.conversationId ? this.messages.length - 1 : index;
|
|
92
|
-
const messageId = this.createId("msg");
|
|
93
|
-
const baseContext = this.buildTextEventContext();
|
|
94
|
-
const content = this.getContentString(message.content);
|
|
95
|
-
aiEventClient.emit("text:message:created", {
|
|
96
|
-
...baseContext,
|
|
97
|
-
messageId,
|
|
98
|
-
role: message.role,
|
|
99
|
-
content,
|
|
100
|
-
toolCalls: message.toolCalls,
|
|
101
|
-
messageIndex,
|
|
102
|
-
timestamp: Date.now()
|
|
103
|
-
});
|
|
104
|
-
if (message.role === "user") {
|
|
105
|
-
aiEventClient.emit("text:message:user", {
|
|
106
|
-
...baseContext,
|
|
107
|
-
messageId,
|
|
108
|
-
role: "user",
|
|
109
|
-
content,
|
|
110
|
-
messageIndex,
|
|
111
|
-
timestamp: Date.now()
|
|
112
|
-
});
|
|
113
|
-
}
|
|
114
|
-
});
|
|
179
|
+
this.middlewareCtx.options = this.eventOptions;
|
|
180
|
+
this.middlewareCtx.toolNames = this.eventToolNames;
|
|
115
181
|
}
|
|
116
|
-
|
|
117
|
-
if (!this.shouldEmitStreamEnd) {
|
|
118
|
-
return;
|
|
119
|
-
}
|
|
120
|
-
const now = Date.now();
|
|
121
|
-
aiEventClient.emit("text:request:completed", {
|
|
122
|
-
...this.buildTextEventContext(),
|
|
123
|
-
content: this.accumulatedContent,
|
|
124
|
-
messageId: this.currentMessageId || void 0,
|
|
125
|
-
finishReason: this.lastFinishReason || void 0,
|
|
126
|
-
usage: this.finishedEvent?.usage,
|
|
127
|
-
duration: now - this.streamStartTime,
|
|
128
|
-
timestamp: now
|
|
129
|
-
});
|
|
130
|
-
}
|
|
131
|
-
beginCycle() {
|
|
182
|
+
async beginCycle() {
|
|
132
183
|
if (this.cyclePhase === "processText") {
|
|
133
|
-
this.beginIteration();
|
|
184
|
+
await this.beginIteration();
|
|
134
185
|
}
|
|
135
186
|
}
|
|
136
187
|
endCycle() {
|
|
@@ -141,27 +192,26 @@ class TextEngine {
|
|
|
141
192
|
this.cyclePhase = "processText";
|
|
142
193
|
this.iterationCount++;
|
|
143
194
|
}
|
|
144
|
-
beginIteration() {
|
|
195
|
+
async beginIteration() {
|
|
145
196
|
this.currentMessageId = this.createId("msg");
|
|
146
197
|
this.accumulatedContent = "";
|
|
147
198
|
this.finishedEvent = null;
|
|
148
|
-
|
|
149
|
-
|
|
150
|
-
|
|
151
|
-
|
|
152
|
-
|
|
153
|
-
content: "",
|
|
154
|
-
timestamp: Date.now()
|
|
199
|
+
this.middlewareCtx.currentMessageId = this.currentMessageId;
|
|
200
|
+
this.middlewareCtx.accumulatedContent = "";
|
|
201
|
+
await this.middlewareRunner.runOnIteration(this.middlewareCtx, {
|
|
202
|
+
iteration: this.iterationCount,
|
|
203
|
+
messageId: this.currentMessageId
|
|
155
204
|
});
|
|
156
205
|
}
|
|
157
206
|
async *streamModelResponse() {
|
|
158
207
|
const { temperature, topP, maxTokens, metadata, modelOptions } = this.params;
|
|
159
|
-
const tools = this.
|
|
160
|
-
const toolsWithJsonSchemas = tools
|
|
208
|
+
const tools = this.tools;
|
|
209
|
+
const toolsWithJsonSchemas = tools.map((tool) => ({
|
|
161
210
|
...tool,
|
|
162
211
|
inputSchema: tool.inputSchema ? convertSchemaToJsonSchema(tool.inputSchema) : void 0,
|
|
163
212
|
outputSchema: tool.outputSchema ? convertSchemaToJsonSchema(tool.outputSchema) : void 0
|
|
164
213
|
}));
|
|
214
|
+
this.middlewareCtx.phase = "modelStream";
|
|
165
215
|
for await (const chunk of this.adapter.chatStream({
|
|
166
216
|
model: this.params.model,
|
|
167
217
|
messages: this.messages,
|
|
@@ -174,12 +224,22 @@ class TextEngine {
|
|
|
174
224
|
modelOptions,
|
|
175
225
|
systemPrompts: this.systemPrompts
|
|
176
226
|
})) {
|
|
177
|
-
if (this.
|
|
227
|
+
if (this.isCancelled()) {
|
|
178
228
|
break;
|
|
179
229
|
}
|
|
180
230
|
this.totalChunkCount++;
|
|
181
|
-
|
|
182
|
-
|
|
231
|
+
const outputChunks = await this.middlewareRunner.runOnChunk(
|
|
232
|
+
this.middlewareCtx,
|
|
233
|
+
chunk
|
|
234
|
+
);
|
|
235
|
+
for (const outputChunk of outputChunks) {
|
|
236
|
+
yield outputChunk;
|
|
237
|
+
this.handleStreamChunk(outputChunk);
|
|
238
|
+
this.middlewareCtx.chunkIndex++;
|
|
239
|
+
}
|
|
240
|
+
if (chunk.type === "RUN_FINISHED" && chunk.usage) {
|
|
241
|
+
await this.middlewareRunner.runOnUsage(this.middlewareCtx, chunk.usage);
|
|
242
|
+
}
|
|
183
243
|
if (this.earlyTermination) {
|
|
184
244
|
break;
|
|
185
245
|
}
|
|
@@ -220,87 +280,25 @@ class TextEngine {
|
|
|
220
280
|
} else {
|
|
221
281
|
this.accumulatedContent += chunk.delta;
|
|
222
282
|
}
|
|
223
|
-
|
|
224
|
-
...this.buildTextEventContext(),
|
|
225
|
-
messageId: this.currentMessageId || void 0,
|
|
226
|
-
content: this.accumulatedContent,
|
|
227
|
-
delta: chunk.delta,
|
|
228
|
-
timestamp: Date.now()
|
|
229
|
-
});
|
|
283
|
+
this.middlewareCtx.accumulatedContent = this.accumulatedContent;
|
|
230
284
|
}
|
|
231
285
|
handleToolCallStartEvent(chunk) {
|
|
232
286
|
this.toolCallManager.addToolCallStartEvent(chunk);
|
|
233
|
-
aiEventClient.emit("text:chunk:tool-call", {
|
|
234
|
-
...this.buildTextEventContext(),
|
|
235
|
-
messageId: this.currentMessageId || void 0,
|
|
236
|
-
toolCallId: chunk.toolCallId,
|
|
237
|
-
toolName: chunk.toolName,
|
|
238
|
-
index: chunk.index ?? 0,
|
|
239
|
-
arguments: "",
|
|
240
|
-
timestamp: Date.now()
|
|
241
|
-
});
|
|
242
287
|
}
|
|
243
288
|
handleToolCallArgsEvent(chunk) {
|
|
244
289
|
this.toolCallManager.addToolCallArgsEvent(chunk);
|
|
245
|
-
aiEventClient.emit("text:chunk:tool-call", {
|
|
246
|
-
...this.buildTextEventContext(),
|
|
247
|
-
messageId: this.currentMessageId || void 0,
|
|
248
|
-
toolCallId: chunk.toolCallId,
|
|
249
|
-
toolName: "",
|
|
250
|
-
index: 0,
|
|
251
|
-
arguments: chunk.delta,
|
|
252
|
-
timestamp: Date.now()
|
|
253
|
-
});
|
|
254
290
|
}
|
|
255
291
|
handleToolCallEndEvent(chunk) {
|
|
256
292
|
this.toolCallManager.completeToolCall(chunk);
|
|
257
|
-
aiEventClient.emit("text:chunk:tool-result", {
|
|
258
|
-
...this.buildTextEventContext(),
|
|
259
|
-
messageId: this.currentMessageId || void 0,
|
|
260
|
-
toolCallId: chunk.toolCallId,
|
|
261
|
-
result: chunk.result || "",
|
|
262
|
-
timestamp: Date.now()
|
|
263
|
-
});
|
|
264
293
|
}
|
|
265
294
|
handleRunFinishedEvent(chunk) {
|
|
266
|
-
aiEventClient.emit("text:chunk:done", {
|
|
267
|
-
...this.buildTextEventContext(),
|
|
268
|
-
messageId: this.currentMessageId || void 0,
|
|
269
|
-
finishReason: chunk.finishReason,
|
|
270
|
-
usage: chunk.usage,
|
|
271
|
-
timestamp: Date.now()
|
|
272
|
-
});
|
|
273
|
-
if (chunk.usage) {
|
|
274
|
-
aiEventClient.emit("text:usage", {
|
|
275
|
-
...this.buildTextEventContext(),
|
|
276
|
-
messageId: this.currentMessageId || void 0,
|
|
277
|
-
usage: chunk.usage,
|
|
278
|
-
timestamp: Date.now()
|
|
279
|
-
});
|
|
280
|
-
}
|
|
281
295
|
this.finishedEvent = chunk;
|
|
282
296
|
this.lastFinishReason = chunk.finishReason;
|
|
283
297
|
}
|
|
284
|
-
handleRunErrorEvent(
|
|
285
|
-
aiEventClient.emit("text:chunk:error", {
|
|
286
|
-
...this.buildTextEventContext(),
|
|
287
|
-
messageId: this.currentMessageId || void 0,
|
|
288
|
-
error: chunk.error.message,
|
|
289
|
-
timestamp: Date.now()
|
|
290
|
-
});
|
|
298
|
+
handleRunErrorEvent(_chunk) {
|
|
291
299
|
this.earlyTermination = true;
|
|
292
|
-
|
|
293
|
-
|
|
294
|
-
handleStepFinishedEvent(chunk) {
|
|
295
|
-
if (chunk.content || chunk.delta) {
|
|
296
|
-
aiEventClient.emit("text:chunk:thinking", {
|
|
297
|
-
...this.buildTextEventContext(),
|
|
298
|
-
messageId: this.currentMessageId || void 0,
|
|
299
|
-
content: chunk.content || "",
|
|
300
|
-
delta: chunk.delta,
|
|
301
|
-
timestamp: Date.now()
|
|
302
|
-
});
|
|
303
|
-
}
|
|
300
|
+
}
|
|
301
|
+
handleStepFinishedEvent(_chunk) {
|
|
304
302
|
}
|
|
305
303
|
async *checkForPendingToolCalls() {
|
|
306
304
|
const pendingToolCalls = this.getPendingToolCallsFromMessages();
|
|
@@ -314,34 +312,65 @@ class TextEngine {
|
|
|
314
312
|
this.tools,
|
|
315
313
|
approvals,
|
|
316
314
|
clientToolResults,
|
|
317
|
-
(eventName, data) => this.createCustomEventChunk(eventName, data)
|
|
315
|
+
(eventName, data) => this.createCustomEventChunk(eventName, data),
|
|
316
|
+
{
|
|
317
|
+
onBeforeToolCall: async (toolCall, tool, args) => {
|
|
318
|
+
const hookCtx = {
|
|
319
|
+
toolCall,
|
|
320
|
+
tool,
|
|
321
|
+
args,
|
|
322
|
+
toolName: toolCall.function.name,
|
|
323
|
+
toolCallId: toolCall.id
|
|
324
|
+
};
|
|
325
|
+
return this.middlewareRunner.runOnBeforeToolCall(
|
|
326
|
+
this.middlewareCtx,
|
|
327
|
+
hookCtx
|
|
328
|
+
);
|
|
329
|
+
},
|
|
330
|
+
onAfterToolCall: async (info) => {
|
|
331
|
+
await this.middlewareRunner.runOnAfterToolCall(
|
|
332
|
+
this.middlewareCtx,
|
|
333
|
+
info
|
|
334
|
+
);
|
|
335
|
+
}
|
|
336
|
+
}
|
|
318
337
|
);
|
|
319
338
|
const executionResult = yield* this.drainToolCallGenerator(generator);
|
|
339
|
+
if (this.isMiddlewareAborted()) {
|
|
340
|
+
this.setToolPhase("stop");
|
|
341
|
+
return "stop";
|
|
342
|
+
}
|
|
343
|
+
await this.middlewareRunner.runOnToolPhaseComplete(this.middlewareCtx, {
|
|
344
|
+
toolCalls: pendingToolCalls,
|
|
345
|
+
results: executionResult.results,
|
|
346
|
+
needsApproval: executionResult.needsApproval,
|
|
347
|
+
needsClientExecution: executionResult.needsClientExecution
|
|
348
|
+
});
|
|
320
349
|
if (executionResult.needsApproval.length > 0 || executionResult.needsClientExecution.length > 0) {
|
|
321
350
|
if (executionResult.results.length > 0) {
|
|
322
|
-
for (const chunk of this.
|
|
351
|
+
for (const chunk of this.buildToolResultChunks(
|
|
323
352
|
executionResult.results,
|
|
324
353
|
finishEvent
|
|
325
354
|
)) {
|
|
326
355
|
yield chunk;
|
|
327
356
|
}
|
|
328
357
|
}
|
|
329
|
-
for (const chunk of this.
|
|
358
|
+
for (const chunk of this.buildApprovalChunks(
|
|
330
359
|
executionResult.needsApproval,
|
|
331
360
|
finishEvent
|
|
332
361
|
)) {
|
|
333
362
|
yield chunk;
|
|
334
363
|
}
|
|
335
|
-
for (const chunk of this.
|
|
364
|
+
for (const chunk of this.buildClientToolChunks(
|
|
336
365
|
executionResult.needsClientExecution,
|
|
337
366
|
finishEvent
|
|
338
367
|
)) {
|
|
339
368
|
yield chunk;
|
|
340
369
|
}
|
|
341
|
-
this.
|
|
370
|
+
this.setToolPhase("wait");
|
|
342
371
|
return "wait";
|
|
343
372
|
}
|
|
344
|
-
const toolResultChunks = this.
|
|
373
|
+
const toolResultChunks = this.buildToolResultChunks(
|
|
345
374
|
executionResult.results,
|
|
346
375
|
finishEvent
|
|
347
376
|
);
|
|
@@ -362,31 +391,64 @@ class TextEngine {
|
|
|
362
391
|
return;
|
|
363
392
|
}
|
|
364
393
|
this.addAssistantToolCallMessage(toolCalls);
|
|
394
|
+
this.middlewareCtx.phase = "beforeTools";
|
|
365
395
|
const { approvals, clientToolResults } = this.collectClientState();
|
|
366
396
|
const generator = executeToolCalls(
|
|
367
397
|
toolCalls,
|
|
368
398
|
this.tools,
|
|
369
399
|
approvals,
|
|
370
400
|
clientToolResults,
|
|
371
|
-
(eventName, data) => this.createCustomEventChunk(eventName, data)
|
|
401
|
+
(eventName, data) => this.createCustomEventChunk(eventName, data),
|
|
402
|
+
{
|
|
403
|
+
onBeforeToolCall: async (toolCall, tool, args) => {
|
|
404
|
+
const hookCtx = {
|
|
405
|
+
toolCall,
|
|
406
|
+
tool,
|
|
407
|
+
args,
|
|
408
|
+
toolName: toolCall.function.name,
|
|
409
|
+
toolCallId: toolCall.id
|
|
410
|
+
};
|
|
411
|
+
return this.middlewareRunner.runOnBeforeToolCall(
|
|
412
|
+
this.middlewareCtx,
|
|
413
|
+
hookCtx
|
|
414
|
+
);
|
|
415
|
+
},
|
|
416
|
+
onAfterToolCall: async (info) => {
|
|
417
|
+
await this.middlewareRunner.runOnAfterToolCall(
|
|
418
|
+
this.middlewareCtx,
|
|
419
|
+
info
|
|
420
|
+
);
|
|
421
|
+
}
|
|
422
|
+
}
|
|
372
423
|
);
|
|
373
424
|
const executionResult = yield* this.drainToolCallGenerator(generator);
|
|
425
|
+
this.middlewareCtx.phase = "afterTools";
|
|
426
|
+
if (this.isMiddlewareAborted()) {
|
|
427
|
+
this.setToolPhase("stop");
|
|
428
|
+
return;
|
|
429
|
+
}
|
|
430
|
+
await this.middlewareRunner.runOnToolPhaseComplete(this.middlewareCtx, {
|
|
431
|
+
toolCalls,
|
|
432
|
+
results: executionResult.results,
|
|
433
|
+
needsApproval: executionResult.needsApproval,
|
|
434
|
+
needsClientExecution: executionResult.needsClientExecution
|
|
435
|
+
});
|
|
374
436
|
if (executionResult.needsApproval.length > 0 || executionResult.needsClientExecution.length > 0) {
|
|
375
437
|
if (executionResult.results.length > 0) {
|
|
376
|
-
for (const chunk of this.
|
|
438
|
+
for (const chunk of this.buildToolResultChunks(
|
|
377
439
|
executionResult.results,
|
|
378
440
|
finishEvent
|
|
379
441
|
)) {
|
|
380
442
|
yield chunk;
|
|
381
443
|
}
|
|
382
444
|
}
|
|
383
|
-
for (const chunk of this.
|
|
445
|
+
for (const chunk of this.buildApprovalChunks(
|
|
384
446
|
executionResult.needsApproval,
|
|
385
447
|
finishEvent
|
|
386
448
|
)) {
|
|
387
449
|
yield chunk;
|
|
388
450
|
}
|
|
389
|
-
for (const chunk of this.
|
|
451
|
+
for (const chunk of this.buildClientToolChunks(
|
|
390
452
|
executionResult.needsClientExecution,
|
|
391
453
|
finishEvent
|
|
392
454
|
)) {
|
|
@@ -395,7 +457,7 @@ class TextEngine {
|
|
|
395
457
|
this.setToolPhase("wait");
|
|
396
458
|
return;
|
|
397
459
|
}
|
|
398
|
-
const toolResultChunks = this.
|
|
460
|
+
const toolResultChunks = this.buildToolResultChunks(
|
|
399
461
|
executionResult.results,
|
|
400
462
|
finishEvent
|
|
401
463
|
);
|
|
@@ -409,7 +471,6 @@ class TextEngine {
|
|
|
409
471
|
return this.finishedEvent?.finishReason === "tool_calls" && this.tools.length > 0 && this.toolCallManager.hasToolCalls();
|
|
410
472
|
}
|
|
411
473
|
addAssistantToolCallMessage(toolCalls) {
|
|
412
|
-
const messageId = this.currentMessageId ?? this.createId("msg");
|
|
413
474
|
this.messages = [
|
|
414
475
|
...this.messages,
|
|
415
476
|
{
|
|
@@ -418,14 +479,6 @@ class TextEngine {
|
|
|
418
479
|
toolCalls
|
|
419
480
|
}
|
|
420
481
|
];
|
|
421
|
-
aiEventClient.emit("text:message:created", {
|
|
422
|
-
...this.buildTextEventContext(),
|
|
423
|
-
messageId,
|
|
424
|
-
role: "assistant",
|
|
425
|
-
content: this.accumulatedContent || "",
|
|
426
|
-
toolCalls,
|
|
427
|
-
timestamp: Date.now()
|
|
428
|
-
});
|
|
429
482
|
}
|
|
430
483
|
/**
|
|
431
484
|
* Extract client state (approvals and client tool results) from original messages.
|
|
@@ -470,18 +523,9 @@ class TextEngine {
|
|
|
470
523
|
}
|
|
471
524
|
return { approvals, clientToolResults };
|
|
472
525
|
}
|
|
473
|
-
|
|
526
|
+
buildApprovalChunks(approvals, finishEvent) {
|
|
474
527
|
const chunks = [];
|
|
475
528
|
for (const approval of approvals) {
|
|
476
|
-
aiEventClient.emit("tools:approval:requested", {
|
|
477
|
-
...this.buildTextEventContext(),
|
|
478
|
-
messageId: this.currentMessageId || void 0,
|
|
479
|
-
toolCallId: approval.toolCallId,
|
|
480
|
-
toolName: approval.toolName,
|
|
481
|
-
input: approval.input,
|
|
482
|
-
approvalId: approval.approvalId,
|
|
483
|
-
timestamp: Date.now()
|
|
484
|
-
});
|
|
485
529
|
chunks.push({
|
|
486
530
|
type: "CUSTOM",
|
|
487
531
|
timestamp: Date.now(),
|
|
@@ -500,17 +544,9 @@ class TextEngine {
|
|
|
500
544
|
}
|
|
501
545
|
return chunks;
|
|
502
546
|
}
|
|
503
|
-
|
|
547
|
+
buildClientToolChunks(clientRequests, finishEvent) {
|
|
504
548
|
const chunks = [];
|
|
505
549
|
for (const clientTool of clientRequests) {
|
|
506
|
-
aiEventClient.emit("tools:input:available", {
|
|
507
|
-
...this.buildTextEventContext(),
|
|
508
|
-
messageId: this.currentMessageId || void 0,
|
|
509
|
-
toolCallId: clientTool.toolCallId,
|
|
510
|
-
toolName: clientTool.toolName,
|
|
511
|
-
input: clientTool.input,
|
|
512
|
-
timestamp: Date.now()
|
|
513
|
-
});
|
|
514
550
|
chunks.push({
|
|
515
551
|
type: "CUSTOM",
|
|
516
552
|
timestamp: Date.now(),
|
|
@@ -525,18 +561,9 @@ class TextEngine {
|
|
|
525
561
|
}
|
|
526
562
|
return chunks;
|
|
527
563
|
}
|
|
528
|
-
|
|
564
|
+
buildToolResultChunks(results, finishEvent) {
|
|
529
565
|
const chunks = [];
|
|
530
566
|
for (const result of results) {
|
|
531
|
-
aiEventClient.emit("tools:call:completed", {
|
|
532
|
-
...this.buildTextEventContext(),
|
|
533
|
-
messageId: this.currentMessageId || void 0,
|
|
534
|
-
toolCallId: result.toolCallId,
|
|
535
|
-
toolName: result.toolName,
|
|
536
|
-
result: result.result,
|
|
537
|
-
duration: result.duration ?? 0,
|
|
538
|
-
timestamp: Date.now()
|
|
539
|
-
});
|
|
540
567
|
const content = JSON.stringify(result.result);
|
|
541
568
|
chunks.push({
|
|
542
569
|
type: "TOOL_CALL_END",
|
|
@@ -554,13 +581,6 @@ class TextEngine {
|
|
|
554
581
|
toolCallId: result.toolCallId
|
|
555
582
|
}
|
|
556
583
|
];
|
|
557
|
-
aiEventClient.emit("text:message:created", {
|
|
558
|
-
...this.buildTextEventContext(),
|
|
559
|
-
messageId: this.createId("msg"),
|
|
560
|
-
role: "tool",
|
|
561
|
-
content,
|
|
562
|
-
timestamp: Date.now()
|
|
563
|
-
});
|
|
564
584
|
}
|
|
565
585
|
return chunks;
|
|
566
586
|
}
|
|
@@ -617,33 +637,44 @@ class TextEngine {
|
|
|
617
637
|
isAborted() {
|
|
618
638
|
return !!this.effectiveSignal?.aborted;
|
|
619
639
|
}
|
|
620
|
-
|
|
640
|
+
isMiddlewareAborted() {
|
|
641
|
+
return !!this.middlewareAbortController?.signal.aborted;
|
|
642
|
+
}
|
|
643
|
+
isCancelled() {
|
|
644
|
+
return this.isAborted() || this.isMiddlewareAborted();
|
|
645
|
+
}
|
|
646
|
+
buildMiddlewareConfig() {
|
|
621
647
|
return {
|
|
622
|
-
|
|
623
|
-
|
|
624
|
-
|
|
625
|
-
|
|
626
|
-
|
|
627
|
-
|
|
628
|
-
|
|
629
|
-
|
|
630
|
-
options: this.eventOptions,
|
|
631
|
-
modelOptions: this.params.modelOptions,
|
|
632
|
-
messageCount: this.initialMessageCount,
|
|
633
|
-
hasTools: this.tools.length > 0,
|
|
634
|
-
streaming: true
|
|
648
|
+
messages: this.messages,
|
|
649
|
+
systemPrompts: [...this.systemPrompts],
|
|
650
|
+
tools: [...this.tools],
|
|
651
|
+
temperature: this.params.temperature,
|
|
652
|
+
topP: this.params.topP,
|
|
653
|
+
maxTokens: this.params.maxTokens,
|
|
654
|
+
metadata: this.params.metadata,
|
|
655
|
+
modelOptions: this.params.modelOptions
|
|
635
656
|
};
|
|
636
657
|
}
|
|
637
|
-
|
|
638
|
-
|
|
639
|
-
|
|
640
|
-
|
|
658
|
+
applyMiddlewareConfig(config) {
|
|
659
|
+
this.messages = config.messages;
|
|
660
|
+
this.systemPrompts = config.systemPrompts;
|
|
661
|
+
this.tools = config.tools;
|
|
662
|
+
this.params = {
|
|
663
|
+
...this.params,
|
|
664
|
+
temperature: config.temperature,
|
|
665
|
+
topP: config.topP,
|
|
666
|
+
maxTokens: config.maxTokens,
|
|
667
|
+
metadata: config.metadata,
|
|
668
|
+
modelOptions: config.modelOptions
|
|
669
|
+
};
|
|
670
|
+
this.middlewareCtx.messages = this.messages;
|
|
671
|
+
this.middlewareCtx.systemPrompts = this.systemPrompts;
|
|
672
|
+
this.middlewareCtx.hasTools = this.tools.length > 0;
|
|
673
|
+
this.middlewareCtx.toolNames = this.tools.map((t) => t.name);
|
|
674
|
+
this.middlewareCtx.modelOptions = config.modelOptions;
|
|
641
675
|
}
|
|
642
676
|
setToolPhase(phase) {
|
|
643
677
|
this.toolPhase = phase;
|
|
644
|
-
if (phase === "wait") {
|
|
645
|
-
this.shouldEmitStreamEnd = false;
|
|
646
|
-
}
|
|
647
678
|
}
|
|
648
679
|
/**
|
|
649
680
|
* Drain an executeToolCalls async generator, yielding any CustomEvent chunks
|
|
@@ -687,11 +718,13 @@ function chat(options) {
|
|
|
687
718
|
);
|
|
688
719
|
}
|
|
689
720
|
async function* runStreamingText(options) {
|
|
690
|
-
const { adapter, ...textOptions } = options;
|
|
721
|
+
const { adapter, middleware, context, ...textOptions } = options;
|
|
691
722
|
const model = adapter.model;
|
|
692
723
|
const engine = new TextEngine({
|
|
693
724
|
adapter,
|
|
694
|
-
params: { ...textOptions, model }
|
|
725
|
+
params: { ...textOptions, model },
|
|
726
|
+
middleware,
|
|
727
|
+
context
|
|
695
728
|
});
|
|
696
729
|
for await (const chunk of engine.run()) {
|
|
697
730
|
yield chunk;
|
|
@@ -704,14 +737,16 @@ function runNonStreamingText(options) {
|
|
|
704
737
|
return streamToText(stream);
|
|
705
738
|
}
|
|
706
739
|
async function runAgenticStructuredOutput(options) {
|
|
707
|
-
const { adapter, outputSchema, ...textOptions } = options;
|
|
740
|
+
const { adapter, outputSchema, middleware, context, ...textOptions } = options;
|
|
708
741
|
const model = adapter.model;
|
|
709
742
|
if (!outputSchema) {
|
|
710
743
|
throw new Error("outputSchema is required for structured output");
|
|
711
744
|
}
|
|
712
745
|
const engine = new TextEngine({
|
|
713
746
|
adapter,
|
|
714
|
-
params: { ...textOptions, model }
|
|
747
|
+
params: { ...textOptions, model },
|
|
748
|
+
middleware,
|
|
749
|
+
context
|
|
715
750
|
});
|
|
716
751
|
for await (const _chunk of engine.run()) {
|
|
717
752
|
}
|