@convex-dev/agent 0.5.0-alpha.1 → 0.6.0-alpha.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/README.md +32 -27
- package/dist/UIMessages.d.ts +46 -0
- package/dist/UIMessages.d.ts.map +1 -0
- package/dist/UIMessages.js +546 -0
- package/dist/UIMessages.js.map +1 -0
- package/dist/client/createTool.d.ts +129 -27
- package/dist/client/createTool.d.ts.map +1 -1
- package/dist/client/createTool.js +66 -12
- package/dist/client/createTool.js.map +1 -1
- package/dist/client/defaultComponent.d.ts +11 -0
- package/dist/client/defaultComponent.d.ts.map +1 -0
- package/dist/client/defaultComponent.js +7 -0
- package/dist/client/defaultComponent.js.map +1 -0
- package/dist/client/definePlaygroundAPI.d.ts +1323 -192
- package/dist/client/definePlaygroundAPI.d.ts.map +1 -1
- package/dist/client/definePlaygroundAPI.js +52 -28
- package/dist/client/definePlaygroundAPI.js.map +1 -1
- package/dist/client/files.d.ts +20 -7
- package/dist/client/files.d.ts.map +1 -1
- package/dist/client/files.js +68 -11
- package/dist/client/files.js.map +1 -1
- package/dist/client/index.d.ts +1056 -965
- package/dist/client/index.d.ts.map +1 -1
- package/dist/client/index.js +242 -748
- package/dist/client/index.js.map +1 -1
- package/dist/client/messages.d.ts +461 -0
- package/dist/client/messages.d.ts.map +1 -0
- package/dist/client/messages.js +106 -0
- package/dist/client/messages.js.map +1 -0
- package/dist/client/mockModel.d.ts +42 -0
- package/dist/client/mockModel.d.ts.map +1 -0
- package/dist/client/mockModel.js +175 -0
- package/dist/client/mockModel.js.map +1 -0
- package/dist/client/saveInputMessages.d.ts +20 -0
- package/dist/client/saveInputMessages.d.ts.map +1 -0
- package/dist/client/saveInputMessages.js +58 -0
- package/dist/client/saveInputMessages.js.map +1 -0
- package/dist/client/search.d.ts +346 -35
- package/dist/client/search.d.ts.map +1 -1
- package/dist/client/search.js +350 -39
- package/dist/client/search.js.map +1 -1
- package/dist/client/start.d.ts +84 -0
- package/dist/client/start.d.ts.map +1 -0
- package/dist/client/start.js +171 -0
- package/dist/client/start.js.map +1 -0
- package/dist/client/streamText.d.ts +46 -0
- package/dist/client/streamText.d.ts.map +1 -0
- package/dist/client/streamText.js +93 -0
- package/dist/client/streamText.js.map +1 -0
- package/dist/client/streaming.d.ts +3705 -32
- package/dist/client/streaming.d.ts.map +1 -1
- package/dist/client/streaming.js +141 -59
- package/dist/client/streaming.js.map +1 -1
- package/dist/client/threads.d.ts +46 -0
- package/dist/client/threads.d.ts.map +1 -0
- package/dist/client/threads.js +49 -0
- package/dist/client/threads.js.map +1 -0
- package/dist/client/types.d.ts +265 -128
- package/dist/client/types.d.ts.map +1 -1
- package/dist/client/utils.d.ts +4 -0
- package/dist/client/utils.d.ts.map +1 -0
- package/dist/client/utils.js +21 -0
- package/dist/client/utils.js.map +1 -0
- package/dist/component/_generated/api.d.ts +24 -2178
- package/dist/component/_generated/api.d.ts.map +1 -1
- package/dist/component/_generated/api.js +10 -1
- package/dist/component/_generated/api.js.map +1 -1
- package/dist/component/_generated/component.d.ts +3119 -0
- package/dist/component/_generated/component.d.ts.map +1 -0
- package/dist/component/_generated/component.js +11 -0
- package/dist/component/_generated/component.js.map +1 -0
- package/dist/component/_generated/dataModel.d.ts +4 -18
- package/dist/component/_generated/dataModel.d.ts.map +1 -0
- package/dist/component/_generated/dataModel.js +11 -0
- package/dist/component/_generated/dataModel.js.map +1 -0
- package/dist/component/_generated/server.d.ts +10 -38
- package/dist/component/_generated/server.d.ts.map +1 -1
- package/dist/component/_generated/server.js +9 -5
- package/dist/component/_generated/server.js.map +1 -1
- package/dist/component/files.d.ts +16 -10
- package/dist/component/files.d.ts.map +1 -1
- package/dist/component/files.js +10 -2
- package/dist/component/files.js.map +1 -1
- package/dist/component/messages.d.ts +2553 -342
- package/dist/component/messages.d.ts.map +1 -1
- package/dist/component/messages.js +387 -154
- package/dist/component/messages.js.map +1 -1
- package/dist/component/schema.d.ts +5697 -3584
- package/dist/component/schema.d.ts.map +1 -1
- package/dist/component/schema.js +18 -41
- package/dist/component/schema.js.map +1 -1
- package/dist/component/streams.d.ts +35 -335
- package/dist/component/streams.d.ts.map +1 -1
- package/dist/component/streams.js +114 -73
- package/dist/component/streams.js.map +1 -1
- package/dist/component/threads.d.ts +16 -16
- package/dist/component/users.d.ts +4 -4
- package/dist/component/vector/index.d.ts +1 -1
- package/dist/component/vector/index.d.ts.map +1 -1
- package/dist/component/vector/index.js +1 -3
- package/dist/component/vector/index.js.map +1 -1
- package/dist/deltas.d.ts +43 -0
- package/dist/deltas.d.ts.map +1 -0
- package/dist/deltas.js +447 -0
- package/dist/deltas.js.map +1 -0
- package/dist/mapping.d.ts +20 -20
- package/dist/mapping.d.ts.map +1 -1
- package/dist/mapping.js +313 -96
- package/dist/mapping.js.map +1 -1
- package/dist/react/SmoothText.d.ts +5 -0
- package/dist/react/SmoothText.d.ts.map +1 -0
- package/dist/react/SmoothText.js +6 -0
- package/dist/react/SmoothText.js.map +1 -0
- package/dist/react/index.d.ts +5 -77
- package/dist/react/index.d.ts.map +1 -1
- package/dist/react/index.js +6 -160
- package/dist/react/index.js.map +1 -1
- package/dist/react/optimisticallySendMessage.d.ts +36 -3
- package/dist/react/optimisticallySendMessage.d.ts.map +1 -1
- package/dist/react/optimisticallySendMessage.js +35 -9
- package/dist/react/optimisticallySendMessage.js.map +1 -1
- package/dist/react/types.d.ts +4 -18
- package/dist/react/types.d.ts.map +1 -1
- package/dist/react/useDeltaStreams.d.ts +10 -0
- package/dist/react/useDeltaStreams.d.ts.map +1 -0
- package/dist/react/useDeltaStreams.js +101 -0
- package/dist/react/useDeltaStreams.js.map +1 -0
- package/dist/react/useSmoothText.d.ts +13 -12
- package/dist/react/useSmoothText.d.ts.map +1 -1
- package/dist/react/useSmoothText.js +32 -15
- package/dist/react/useSmoothText.js.map +1 -1
- package/dist/react/useStreamingUIMessages.d.ts +22 -0
- package/dist/react/useStreamingUIMessages.d.ts.map +1 -0
- package/dist/react/useStreamingUIMessages.js +92 -0
- package/dist/react/useStreamingUIMessages.js.map +1 -0
- package/dist/react/useThreadMessages.d.ts +104 -0
- package/dist/react/useThreadMessages.d.ts.map +1 -0
- package/dist/react/useThreadMessages.js +148 -0
- package/dist/react/useThreadMessages.js.map +1 -0
- package/dist/react/useUIMessages.d.ts +96 -0
- package/dist/react/useUIMessages.d.ts.map +1 -0
- package/dist/react/useUIMessages.js +108 -0
- package/dist/react/useUIMessages.js.map +1 -0
- package/dist/shared.d.ts +20 -4
- package/dist/shared.d.ts.map +1 -1
- package/dist/shared.js +45 -8
- package/dist/shared.js.map +1 -1
- package/dist/validators.d.ts +22981 -5666
- package/dist/validators.d.ts.map +1 -1
- package/dist/validators.js +245 -137
- package/dist/validators.js.map +1 -1
- package/package.json +98 -50
- package/src/UIMessages.combineUIMessages.test.ts +239 -0
- package/src/UIMessages.test.ts +273 -0
- package/src/UIMessages.ts +739 -0
- package/src/client/createTool.ts +293 -76
- package/src/client/defaultComponent.ts +17 -0
- package/src/client/definePlaygroundAPI.ts +67 -31
- package/src/client/files.ts +100 -20
- package/src/client/index.test.ts +40 -85
- package/src/client/index.ts +520 -1290
- package/src/client/messages.ts +237 -0
- package/src/client/mockModel.ts +245 -0
- package/src/client/saveInputMessages.test.ts +583 -0
- package/src/client/saveInputMessages.ts +101 -0
- package/src/client/search.test.ts +1207 -0
- package/src/client/search.ts +577 -70
- package/src/client/start.ts +310 -0
- package/src/client/streamText.ts +163 -0
- package/src/client/streaming.test.ts +186 -0
- package/src/client/streaming.ts +219 -97
- package/src/client/threads.ts +83 -0
- package/src/client/types.ts +368 -219
- package/src/client/utils.ts +27 -0
- package/src/component/_generated/api.ts +64 -0
- package/src/component/_generated/component.ts +4913 -0
- package/src/component/_generated/{server.d.ts → server.ts} +33 -21
- package/src/component/files.ts +11 -2
- package/src/component/messages.test.ts +195 -51
- package/src/component/messages.ts +490 -201
- package/src/component/schema.ts +20 -46
- package/src/component/setup.test.ts +7 -0
- package/src/component/streams.ts +184 -83
- package/src/component/users.test.ts +0 -1
- package/src/component/vector/index.ts +1 -3
- package/src/deltas.test.ts +626 -0
- package/src/deltas.ts +570 -0
- package/src/fromUIMessages.test.ts +497 -0
- package/src/mapping.test.ts +103 -6
- package/src/mapping.ts +422 -161
- package/src/react/SmoothText.tsx +9 -0
- package/src/react/index.ts +10 -230
- package/src/react/optimisticallySendMessage.ts +55 -12
- package/src/react/types.ts +6 -39
- package/src/react/useDeltaStreams.ts +154 -0
- package/src/react/useSmoothText.ts +56 -36
- package/src/react/useStreamingUIMessages.ts +143 -0
- package/src/react/useThreadMessages.ts +262 -0
- package/src/react/useUIMessages.test.ts +255 -0
- package/src/react/useUIMessages.ts +195 -0
- package/src/shared.ts +88 -12
- package/src/test.ts +18 -0
- package/src/toUIMessages.test.ts +1269 -0
- package/src/validators.test.ts +18 -19
- package/src/validators.ts +325 -185
- package/dist/client/_generated/_ignore.d.ts +0 -1
- package/dist/client/_generated/_ignore.d.ts.map +0 -1
- package/dist/client/_generated/_ignore.js +0 -3
- package/dist/client/_generated/_ignore.js.map +0 -1
- package/dist/client/listMessages.d.ts +0 -22
- package/dist/client/listMessages.d.ts.map +0 -1
- package/dist/client/listMessages.js +0 -25
- package/dist/client/listMessages.js.map +0 -1
- package/dist/package.json +0 -3
- package/dist/react/deltas.d.ts +0 -26
- package/dist/react/deltas.d.ts.map +0 -1
- package/dist/react/deltas.js +0 -384
- package/dist/react/deltas.js.map +0 -1
- package/dist/react/toUIMessages.d.ts +0 -15
- package/dist/react/toUIMessages.d.ts.map +0 -1
- package/dist/react/toUIMessages.js +0 -211
- package/dist/react/toUIMessages.js.map +0 -1
- package/src/client/listMessages.ts +0 -38
- package/src/component/_generated/api.d.ts +0 -2202
- package/src/component/_generated/api.js +0 -23
- package/src/component/_generated/server.js +0 -90
- package/src/node_modules/.vite/vitest/da39a3ee5e6b4b0d3255bfef95601890afd80709/results.json +0 -1
- package/src/react/deltas.test.ts +0 -315
- package/src/react/deltas.ts +0 -478
- package/src/react/toUIMessages.test.ts +0 -420
- package/src/react/toUIMessages.ts +0 -253
- package/src/vitest.config.ts +0 -7
- /package/src/component/_generated/{dataModel.d.ts → dataModel.ts} +0 -0
|
@@ -0,0 +1,237 @@
|
|
|
1
|
+
import type { ModelMessage } from "ai";
|
|
2
|
+
import type { PaginationOptions, PaginationResult } from "convex/server";
|
|
3
|
+
import type { MessageDoc } from "../validators.js";
|
|
4
|
+
import { validateVectorDimension } from "../component/vector/tables.js";
|
|
5
|
+
import {
|
|
6
|
+
vMessageWithMetadata,
|
|
7
|
+
type Message,
|
|
8
|
+
type MessageEmbeddings,
|
|
9
|
+
type MessageEmbeddingsWithDimension,
|
|
10
|
+
type MessageStatus,
|
|
11
|
+
type MessageWithMetadata,
|
|
12
|
+
} from "../validators.js";
|
|
13
|
+
import { serializeMessage } from "../mapping.js";
|
|
14
|
+
import { toUIMessages, type UIMessage } from "../UIMessages.js";
|
|
15
|
+
import type {
|
|
16
|
+
AgentComponent,
|
|
17
|
+
MutationCtx,
|
|
18
|
+
QueryCtx,
|
|
19
|
+
ActionCtx,
|
|
20
|
+
} from "./types.js";
|
|
21
|
+
import { parse } from "convex-helpers/validators";
|
|
22
|
+
|
|
23
|
+
/**
|
|
24
|
+
* List messages from a thread.
|
|
25
|
+
* @param ctx A ctx object from a query, mutation, or action.
|
|
26
|
+
* @param component The agent component, usually `components.agent`.
|
|
27
|
+
* @param args.threadId The thread to list messages from.
|
|
28
|
+
* @param args.paginationOpts Pagination options (e.g. via usePaginatedQuery).
|
|
29
|
+
* @param args.excludeToolMessages Whether to exclude tool messages.
|
|
30
|
+
* False by default.
|
|
31
|
+
* @param args.statuses What statuses to include. All by default.
|
|
32
|
+
* @returns The MessageDoc's in a format compatible with usePaginatedQuery.
|
|
33
|
+
*/
|
|
34
|
+
export async function listMessages(
|
|
35
|
+
ctx: QueryCtx | MutationCtx | ActionCtx,
|
|
36
|
+
component: AgentComponent,
|
|
37
|
+
{
|
|
38
|
+
threadId,
|
|
39
|
+
paginationOpts,
|
|
40
|
+
excludeToolMessages,
|
|
41
|
+
statuses,
|
|
42
|
+
}: {
|
|
43
|
+
threadId: string;
|
|
44
|
+
paginationOpts: PaginationOptions;
|
|
45
|
+
excludeToolMessages?: boolean;
|
|
46
|
+
statuses?: MessageStatus[];
|
|
47
|
+
},
|
|
48
|
+
): Promise<PaginationResult<MessageDoc>> {
|
|
49
|
+
if (paginationOpts.numItems === 0) {
|
|
50
|
+
return {
|
|
51
|
+
page: [],
|
|
52
|
+
isDone: true,
|
|
53
|
+
continueCursor: paginationOpts.cursor ?? "",
|
|
54
|
+
};
|
|
55
|
+
}
|
|
56
|
+
return ctx.runQuery(component.messages.listMessagesByThreadId, {
|
|
57
|
+
order: "desc",
|
|
58
|
+
threadId,
|
|
59
|
+
paginationOpts,
|
|
60
|
+
excludeToolMessages,
|
|
61
|
+
statuses,
|
|
62
|
+
});
|
|
63
|
+
}
|
|
64
|
+
|
|
65
|
+
export async function listUIMessages(
|
|
66
|
+
ctx: QueryCtx | MutationCtx | ActionCtx,
|
|
67
|
+
component: AgentComponent,
|
|
68
|
+
args: {
|
|
69
|
+
threadId: string;
|
|
70
|
+
paginationOpts: PaginationOptions;
|
|
71
|
+
},
|
|
72
|
+
): Promise<PaginationResult<UIMessage>> {
|
|
73
|
+
const result = await listMessages(ctx, component, args);
|
|
74
|
+
return { ...result, page: toUIMessages(result.page) };
|
|
75
|
+
}
|
|
76
|
+
|
|
77
|
+
export type SaveMessagesArgs = {
|
|
78
|
+
threadId: string;
|
|
79
|
+
userId?: string | null;
|
|
80
|
+
/**
|
|
81
|
+
* The message that these messages are in response to. They will be
|
|
82
|
+
* the same "order" as this message, at increasing stepOrder(s).
|
|
83
|
+
*/
|
|
84
|
+
promptMessageId?: string;
|
|
85
|
+
/**
|
|
86
|
+
* The messages to save.
|
|
87
|
+
*/
|
|
88
|
+
messages: (ModelMessage | Message)[];
|
|
89
|
+
/**
|
|
90
|
+
* Metadata to save with the messages. Each element corresponds to the
|
|
91
|
+
* message at the same index.
|
|
92
|
+
*/
|
|
93
|
+
metadata?: Omit<MessageWithMetadata, "message">[];
|
|
94
|
+
/**
|
|
95
|
+
* If true, it will fail any pending steps.
|
|
96
|
+
* Defaults to false.
|
|
97
|
+
*/
|
|
98
|
+
failPendingSteps?: boolean;
|
|
99
|
+
/**
|
|
100
|
+
* The embeddings to save with the messages.
|
|
101
|
+
*/
|
|
102
|
+
embeddings?: MessageEmbeddings;
|
|
103
|
+
/**
|
|
104
|
+
* A pending message ID to replace when adding messages.
|
|
105
|
+
*/
|
|
106
|
+
pendingMessageId?: string;
|
|
107
|
+
};
|
|
108
|
+
|
|
109
|
+
/**
|
|
110
|
+
* Explicitly save messages associated with the thread (& user if provided)
|
|
111
|
+
*/
|
|
112
|
+
export async function saveMessages(
|
|
113
|
+
ctx: MutationCtx | ActionCtx,
|
|
114
|
+
component: AgentComponent,
|
|
115
|
+
args: SaveMessagesArgs & {
|
|
116
|
+
/**
|
|
117
|
+
* The agent name to associate with the messages.
|
|
118
|
+
*/
|
|
119
|
+
agentName?: string;
|
|
120
|
+
},
|
|
121
|
+
): Promise<{ messages: MessageDoc[] }> {
|
|
122
|
+
let embeddings: MessageEmbeddingsWithDimension | undefined;
|
|
123
|
+
if (args.embeddings) {
|
|
124
|
+
const dimension = args.embeddings.vectors.find((v) => v !== null)?.length;
|
|
125
|
+
if (dimension) {
|
|
126
|
+
validateVectorDimension(dimension);
|
|
127
|
+
embeddings = {
|
|
128
|
+
model: args.embeddings.model,
|
|
129
|
+
dimension,
|
|
130
|
+
vectors: args.embeddings.vectors,
|
|
131
|
+
};
|
|
132
|
+
}
|
|
133
|
+
}
|
|
134
|
+
const result = await ctx.runMutation(component.messages.addMessages, {
|
|
135
|
+
threadId: args.threadId,
|
|
136
|
+
userId: args.userId ?? undefined,
|
|
137
|
+
agentName: args.agentName,
|
|
138
|
+
promptMessageId: args.promptMessageId,
|
|
139
|
+
pendingMessageId: args.pendingMessageId,
|
|
140
|
+
embeddings,
|
|
141
|
+
messages: await Promise.all(
|
|
142
|
+
args.messages.map(async (m, i) => {
|
|
143
|
+
const { message, fileIds } = await serializeMessage(ctx, component, m);
|
|
144
|
+
const base = args.metadata?.[i];
|
|
145
|
+
const allFileIds = [...(base?.fileIds ?? [])];
|
|
146
|
+
if (fileIds) allFileIds.push(...fileIds);
|
|
147
|
+
|
|
148
|
+
return parse(vMessageWithMetadata, {
|
|
149
|
+
...base,
|
|
150
|
+
message,
|
|
151
|
+
...(allFileIds.length > 0 ? { fileIds: allFileIds } : {}),
|
|
152
|
+
});
|
|
153
|
+
}),
|
|
154
|
+
),
|
|
155
|
+
failPendingSteps: args.failPendingSteps ?? false,
|
|
156
|
+
});
|
|
157
|
+
return { messages: result.messages };
|
|
158
|
+
}
|
|
159
|
+
|
|
160
|
+
export type SaveMessageArgs = {
|
|
161
|
+
threadId: string;
|
|
162
|
+
userId?: string | null;
|
|
163
|
+
/**
|
|
164
|
+
* The message that these messages are in response to. They will be
|
|
165
|
+
* the same "order" as this message, at increasing stepOrder(s).
|
|
166
|
+
*/
|
|
167
|
+
promptMessageId?: string;
|
|
168
|
+
/**
|
|
169
|
+
* Metadata to save with the messages. Each element corresponds to the
|
|
170
|
+
* message at the same index.
|
|
171
|
+
*/
|
|
172
|
+
metadata?: Omit<MessageWithMetadata, "message">;
|
|
173
|
+
/**
|
|
174
|
+
* The embedding to save with the message.
|
|
175
|
+
*/
|
|
176
|
+
embedding?: { vector: number[]; model: string };
|
|
177
|
+
/**
|
|
178
|
+
* A pending message ID to replace with this message.
|
|
179
|
+
*/
|
|
180
|
+
pendingMessageId?: string;
|
|
181
|
+
} & (
|
|
182
|
+
| {
|
|
183
|
+
prompt?: undefined;
|
|
184
|
+
/**
|
|
185
|
+
* The message to save.
|
|
186
|
+
*/
|
|
187
|
+
message: ModelMessage | Message;
|
|
188
|
+
}
|
|
189
|
+
| {
|
|
190
|
+
/*
|
|
191
|
+
* The prompt to save with the message.
|
|
192
|
+
*/
|
|
193
|
+
prompt: string;
|
|
194
|
+
message?: undefined;
|
|
195
|
+
}
|
|
196
|
+
);
|
|
197
|
+
|
|
198
|
+
/**
|
|
199
|
+
* Save a message to the thread.
|
|
200
|
+
* @param ctx A ctx object from a mutation or action.
|
|
201
|
+
* @param args The message and what to associate it with (user / thread)
|
|
202
|
+
* You can pass extra metadata alongside the message, e.g. associated fileIds.
|
|
203
|
+
* @returns The messageId of the saved message.
|
|
204
|
+
*/
|
|
205
|
+
export async function saveMessage(
|
|
206
|
+
ctx: MutationCtx | ActionCtx,
|
|
207
|
+
component: AgentComponent,
|
|
208
|
+
args: SaveMessageArgs & {
|
|
209
|
+
/**
|
|
210
|
+
* The agent name to associate with the message.
|
|
211
|
+
*/
|
|
212
|
+
agentName?: string;
|
|
213
|
+
},
|
|
214
|
+
) {
|
|
215
|
+
let embeddings: { vectors: number[][]; model: string } | undefined;
|
|
216
|
+
if (args.embedding && args.embedding.vector) {
|
|
217
|
+
embeddings = {
|
|
218
|
+
model: args.embedding.model,
|
|
219
|
+
vectors: [args.embedding.vector],
|
|
220
|
+
};
|
|
221
|
+
}
|
|
222
|
+
const { messages } = await saveMessages(ctx, component, {
|
|
223
|
+
threadId: args.threadId,
|
|
224
|
+
userId: args.userId ?? undefined,
|
|
225
|
+
agentName: args.agentName,
|
|
226
|
+
promptMessageId: args.promptMessageId,
|
|
227
|
+
pendingMessageId: args.pendingMessageId,
|
|
228
|
+
messages:
|
|
229
|
+
args.prompt !== undefined
|
|
230
|
+
? [{ role: "user", content: args.prompt }]
|
|
231
|
+
: [args.message],
|
|
232
|
+
metadata: args.metadata ? [args.metadata] : undefined,
|
|
233
|
+
embeddings,
|
|
234
|
+
});
|
|
235
|
+
const message = messages.at(-1)!;
|
|
236
|
+
return { messageId: message._id, message };
|
|
237
|
+
}
|
|
@@ -0,0 +1,245 @@
|
|
|
1
|
+
import type {
|
|
2
|
+
LanguageModelV3,
|
|
3
|
+
LanguageModelV3Content,
|
|
4
|
+
LanguageModelV3StreamPart,
|
|
5
|
+
} from "@ai-sdk/provider";
|
|
6
|
+
import { simulateReadableStream, type ProviderMetadata } from "ai";
|
|
7
|
+
import { assert, pick } from "convex-helpers";
|
|
8
|
+
|
|
9
|
+
export const DEFAULT_TEXT = `
|
|
10
|
+
A A A A A A A A A A A A A A A
|
|
11
|
+
B B B B B B B B B B B B B B B
|
|
12
|
+
C C C C C C C C C C C C C C C
|
|
13
|
+
D D D D D D D D D D D D D D D
|
|
14
|
+
`;
|
|
15
|
+
const DEFAULT_USAGE = {
|
|
16
|
+
outputTokens: 10,
|
|
17
|
+
inputTokens: 3,
|
|
18
|
+
totalTokens: 13,
|
|
19
|
+
inputTokenDetails: undefined,
|
|
20
|
+
outputTokenDetails: undefined,
|
|
21
|
+
};
|
|
22
|
+
|
|
23
|
+
export type MockModelArgs = {
|
|
24
|
+
provider?: LanguageModelV3["provider"];
|
|
25
|
+
modelId?: LanguageModelV3["modelId"];
|
|
26
|
+
supportedUrls?:
|
|
27
|
+
| LanguageModelV3["supportedUrls"]
|
|
28
|
+
| (() => LanguageModelV3["supportedUrls"]);
|
|
29
|
+
chunkDelayInMs?: number;
|
|
30
|
+
initialDelayInMs?: number;
|
|
31
|
+
/** A list of the responses for multiple steps.
|
|
32
|
+
* For tool calls, the first list would include a tool call part,
|
|
33
|
+
* then the next list would be after the tool response or another tool call.
|
|
34
|
+
* Tool responses come from actual tool calls!
|
|
35
|
+
*/
|
|
36
|
+
contentSteps?: LanguageModelV3Content[][];
|
|
37
|
+
/** A single list of content responded from each step.
|
|
38
|
+
* Provide contentSteps instead if you want to do multi-step responses with
|
|
39
|
+
* tool calls.
|
|
40
|
+
*/
|
|
41
|
+
content?: LanguageModelV3Content[];
|
|
42
|
+
// provide either content, contentResponses or doGenerate & doStream
|
|
43
|
+
doGenerate?: LanguageModelV3["doGenerate"];
|
|
44
|
+
doStream?: LanguageModelV3["doStream"];
|
|
45
|
+
providerMetadata?: ProviderMetadata;
|
|
46
|
+
fail?:
|
|
47
|
+
| boolean
|
|
48
|
+
| {
|
|
49
|
+
probability?: number;
|
|
50
|
+
error?: string;
|
|
51
|
+
};
|
|
52
|
+
};
|
|
53
|
+
|
|
54
|
+
function atMostOneOf(...args: unknown[]) {
|
|
55
|
+
return args.filter(Boolean).length <= 1;
|
|
56
|
+
}
|
|
57
|
+
|
|
58
|
+
export function mockModel(args?: MockModelArgs): LanguageModelV3 {
|
|
59
|
+
return new MockLanguageModel(args ?? {});
|
|
60
|
+
}
|
|
61
|
+
|
|
62
|
+
export class MockLanguageModel implements LanguageModelV3 {
|
|
63
|
+
readonly specificationVersion = "v3";
|
|
64
|
+
|
|
65
|
+
private _supportedUrls: () => LanguageModelV3["supportedUrls"];
|
|
66
|
+
|
|
67
|
+
readonly provider: LanguageModelV3["provider"];
|
|
68
|
+
readonly modelId: LanguageModelV3["modelId"];
|
|
69
|
+
|
|
70
|
+
doGenerate: LanguageModelV3["doGenerate"];
|
|
71
|
+
doStream: LanguageModelV3["doStream"];
|
|
72
|
+
|
|
73
|
+
doGenerateCalls: Parameters<LanguageModelV3["doGenerate"]>[0][] = [];
|
|
74
|
+
doStreamCalls: Parameters<LanguageModelV3["doStream"]>[0][] = [];
|
|
75
|
+
|
|
76
|
+
constructor(args: MockModelArgs) {
|
|
77
|
+
assert(
|
|
78
|
+
atMostOneOf(
|
|
79
|
+
args.content,
|
|
80
|
+
args.contentSteps,
|
|
81
|
+
args.doGenerate && args.doStream,
|
|
82
|
+
),
|
|
83
|
+
"Expected only one of content, contentSteps, or doGenerate and doStream",
|
|
84
|
+
);
|
|
85
|
+
this.provider = args.provider || "mock-provider";
|
|
86
|
+
this.modelId = args.modelId || "mock-model-id";
|
|
87
|
+
const {
|
|
88
|
+
content = [{ type: "text", text: DEFAULT_TEXT }],
|
|
89
|
+
contentSteps = [content],
|
|
90
|
+
chunkDelayInMs = 0,
|
|
91
|
+
initialDelayInMs = 0,
|
|
92
|
+
supportedUrls = {},
|
|
93
|
+
} = args;
|
|
94
|
+
const fail =
|
|
95
|
+
args.fail &&
|
|
96
|
+
(args.fail === true ||
|
|
97
|
+
!args.fail.probability ||
|
|
98
|
+
Math.random() < args.fail.probability);
|
|
99
|
+
const error =
|
|
100
|
+
(typeof args.fail === "object" && args.fail.error) ||
|
|
101
|
+
"Mock error message";
|
|
102
|
+
const metadata = pick(args, ["providerMetadata"]);
|
|
103
|
+
|
|
104
|
+
const chunkResponses: LanguageModelV3StreamPart[][] = contentSteps.map(
|
|
105
|
+
(content) => {
|
|
106
|
+
const chunks: LanguageModelV3StreamPart[] = [
|
|
107
|
+
{ type: "stream-start", warnings: [] },
|
|
108
|
+
];
|
|
109
|
+
chunks.push(
|
|
110
|
+
...content.flatMap((c, ci): LanguageModelV3StreamPart[] => {
|
|
111
|
+
if (c.type !== "text" && c.type !== "reasoning") {
|
|
112
|
+
return [c];
|
|
113
|
+
}
|
|
114
|
+
const metadata = pick(c, ["providerMetadata"]);
|
|
115
|
+
const deltas = c.text.split(" ");
|
|
116
|
+
const parts: LanguageModelV3StreamPart[] = [];
|
|
117
|
+
if (c.type === "reasoning") {
|
|
118
|
+
parts.push({
|
|
119
|
+
type: "reasoning-start",
|
|
120
|
+
id: `reasoning-${ci}`,
|
|
121
|
+
...metadata,
|
|
122
|
+
});
|
|
123
|
+
parts.push(
|
|
124
|
+
...deltas.map(
|
|
125
|
+
(delta, di) =>
|
|
126
|
+
({
|
|
127
|
+
type: "reasoning-delta",
|
|
128
|
+
delta: (di ? " " : "") + delta,
|
|
129
|
+
id: `reasoning-${ci}`,
|
|
130
|
+
...metadata,
|
|
131
|
+
}) satisfies LanguageModelV3StreamPart,
|
|
132
|
+
),
|
|
133
|
+
);
|
|
134
|
+
parts.push({
|
|
135
|
+
type: "reasoning-end",
|
|
136
|
+
id: `reasoning-${ci}`,
|
|
137
|
+
...metadata,
|
|
138
|
+
});
|
|
139
|
+
} else if (c.type === "text") {
|
|
140
|
+
parts.push({
|
|
141
|
+
type: "text-start",
|
|
142
|
+
id: `txt-${ci}`,
|
|
143
|
+
...metadata,
|
|
144
|
+
});
|
|
145
|
+
parts.push(
|
|
146
|
+
...deltas.map(
|
|
147
|
+
(delta, di) =>
|
|
148
|
+
({
|
|
149
|
+
type: "text-delta",
|
|
150
|
+
delta: (di ? " " : "") + delta,
|
|
151
|
+
id: `txt-${ci}`,
|
|
152
|
+
...metadata,
|
|
153
|
+
}) satisfies LanguageModelV3StreamPart,
|
|
154
|
+
),
|
|
155
|
+
);
|
|
156
|
+
parts.push({
|
|
157
|
+
type: "text-end",
|
|
158
|
+
id: `txt-${ci}`,
|
|
159
|
+
...metadata,
|
|
160
|
+
});
|
|
161
|
+
}
|
|
162
|
+
return parts;
|
|
163
|
+
}),
|
|
164
|
+
);
|
|
165
|
+
if (fail) {
|
|
166
|
+
chunks.push({
|
|
167
|
+
type: "error",
|
|
168
|
+
error,
|
|
169
|
+
});
|
|
170
|
+
}
|
|
171
|
+
chunks.push({
|
|
172
|
+
type: "finish",
|
|
173
|
+
finishReason: fail ? "error" : "stop",
|
|
174
|
+
usage: DEFAULT_USAGE,
|
|
175
|
+
...(metadata as any),
|
|
176
|
+
});
|
|
177
|
+
return chunks;
|
|
178
|
+
},
|
|
179
|
+
);
|
|
180
|
+
let callIndex = 0;
|
|
181
|
+
this.doGenerate = async (options) => {
|
|
182
|
+
this.doGenerateCalls.push(options);
|
|
183
|
+
|
|
184
|
+
if (fail) {
|
|
185
|
+
throw new Error(error);
|
|
186
|
+
}
|
|
187
|
+
if (typeof args.doGenerate === "function") {
|
|
188
|
+
return args.doGenerate(options);
|
|
189
|
+
} else if (Array.isArray(args.doGenerate)) {
|
|
190
|
+
return args.doGenerate[this.doGenerateCalls.length];
|
|
191
|
+
} else if (contentSteps.length) {
|
|
192
|
+
const result = {
|
|
193
|
+
content: contentSteps[callIndex % contentSteps.length],
|
|
194
|
+
finishReason: "stop" as const,
|
|
195
|
+
usage: DEFAULT_USAGE,
|
|
196
|
+
...(metadata as any),
|
|
197
|
+
warnings: [],
|
|
198
|
+
};
|
|
199
|
+
callIndex++;
|
|
200
|
+
return result;
|
|
201
|
+
} else {
|
|
202
|
+
throw new Error("Unexpected: no content or doGenerate");
|
|
203
|
+
}
|
|
204
|
+
};
|
|
205
|
+
this.doStream = async (options) => {
|
|
206
|
+
this.doStreamCalls.push(options);
|
|
207
|
+
|
|
208
|
+
if (typeof args.doStream === "function") {
|
|
209
|
+
return args.doStream(options);
|
|
210
|
+
} else if (Array.isArray(args.doStream)) {
|
|
211
|
+
return args.doStream[this.doStreamCalls.length];
|
|
212
|
+
} else if (contentSteps) {
|
|
213
|
+
const stream = simulateReadableStream({
|
|
214
|
+
chunks: chunkResponses[callIndex % chunkResponses.length],
|
|
215
|
+
initialDelayInMs,
|
|
216
|
+
chunkDelayInMs,
|
|
217
|
+
});
|
|
218
|
+
callIndex++;
|
|
219
|
+
|
|
220
|
+
if (options.abortSignal) {
|
|
221
|
+
options.abortSignal.addEventListener("abort", () => {
|
|
222
|
+
console.warn("abortSignal in mock model not supported");
|
|
223
|
+
});
|
|
224
|
+
}
|
|
225
|
+
return {
|
|
226
|
+
stream,
|
|
227
|
+
request: { body: {} },
|
|
228
|
+
response: { headers: {} },
|
|
229
|
+
};
|
|
230
|
+
} else if (args.doStream) {
|
|
231
|
+
return args.doStream;
|
|
232
|
+
} else {
|
|
233
|
+
throw new Error("Provide either content or doStream");
|
|
234
|
+
}
|
|
235
|
+
};
|
|
236
|
+
this._supportedUrls =
|
|
237
|
+
typeof supportedUrls === "function"
|
|
238
|
+
? supportedUrls
|
|
239
|
+
: async () => supportedUrls;
|
|
240
|
+
}
|
|
241
|
+
|
|
242
|
+
get supportedUrls() {
|
|
243
|
+
return this._supportedUrls();
|
|
244
|
+
}
|
|
245
|
+
}
|