@assistant-ui/react-google-adk 0.0.35 → 0.0.36
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 +12 -2
- package/dist/AdkClient.d.ts +4 -0
- package/dist/AdkClient.d.ts.map +1 -1
- package/dist/AdkClient.js +12 -7
- package/dist/AdkClient.js.map +1 -1
- package/dist/AdkSessionAdapter.d.ts.map +1 -1
- package/dist/AdkSessionAdapter.js +7 -7
- package/dist/AdkSessionAdapter.js.map +1 -1
- package/dist/AdkThreadController.d.ts +15 -0
- package/dist/AdkThreadController.d.ts.map +1 -0
- package/dist/AdkThreadController.js +35 -0
- package/dist/AdkThreadController.js.map +1 -0
- package/dist/adkThreadState.d.ts +54 -0
- package/dist/adkThreadState.d.ts.map +1 -0
- package/dist/adkThreadState.js +93 -0
- package/dist/adkThreadState.js.map +1 -0
- package/dist/convertToAdkMessages.js +1 -1
- package/dist/convertToAdkMessages.js.map +1 -1
- package/dist/sdkIdentity.js +1 -1
- package/dist/server/createAdkApiRoute.d.ts +37 -6
- package/dist/server/createAdkApiRoute.d.ts.map +1 -1
- package/dist/server/createAdkApiRoute.js +55 -5
- package/dist/server/createAdkApiRoute.js.map +1 -1
- package/dist/server/parseAdkRequest.d.ts +4 -1
- package/dist/server/parseAdkRequest.d.ts.map +1 -1
- package/dist/server/parseAdkRequest.js +5 -1
- package/dist/server/parseAdkRequest.js.map +1 -1
- package/dist/useAdkMessages.d.ts +9 -7
- package/dist/useAdkMessages.d.ts.map +1 -1
- package/dist/useAdkMessages.js +55 -77
- package/dist/useAdkMessages.js.map +1 -1
- package/dist/useAdkRuntime.d.ts +7 -1
- package/dist/useAdkRuntime.d.ts.map +1 -1
- package/dist/useAdkRuntime.js +134 -56
- package/dist/useAdkRuntime.js.map +1 -1
- package/package.json +5 -5
- package/src/AdkClient.test.ts +78 -2
- package/src/AdkClient.ts +24 -6
- package/src/AdkSessionAdapter.ts +1 -1
- package/src/AdkThreadController.test.ts +90 -0
- package/src/AdkThreadController.ts +45 -0
- package/src/adkThreadState.test.ts +207 -0
- package/src/adkThreadState.ts +124 -0
- package/src/convertToAdkMessages.test.ts +19 -0
- package/src/convertToAdkMessages.ts +1 -1
- package/src/hooks.test.tsx +1 -0
- package/src/server/createAdkApiRoute.controls.test.ts +66 -0
- package/src/server/createAdkApiRoute.test.ts +282 -0
- package/src/server/createAdkApiRoute.ts +119 -11
- package/src/server/parseAdkRequest.test.ts +11 -3
- package/src/server/parseAdkRequest.ts +7 -1
- package/src/useAdkMessages.test.ts +1 -0
- package/src/useAdkMessages.ts +61 -96
- package/src/useAdkRuntime.cancellation.test.tsx +4 -3
- package/src/useAdkRuntime.cloud-options.test.tsx +59 -0
- package/src/useAdkRuntime.refetch.test.tsx +548 -4
- package/src/useAdkRuntime.replacement.test.tsx +718 -1
- package/src/useAdkRuntime.ts +169 -73
- package/src/useAdkRuntimeApproval.test.tsx +87 -1
- package/dist/raceWithAbortSignal.d.ts +0 -2
- package/dist/raceWithAbortSignal.d.ts.map +0 -1
- package/dist/raceWithAbortSignal.js +0 -45
- package/dist/raceWithAbortSignal.js.map +0 -1
- package/src/raceWithAbortSignal.test.ts +0 -73
- package/src/raceWithAbortSignal.ts +0 -48
package/src/hooks.test.tsx
CHANGED
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
import { describe, expect, it, vi } from "vitest";
|
|
2
|
+
import { createAdkApiRoute } from "./createAdkApiRoute";
|
|
3
|
+
|
|
4
|
+
const makeRequest = (body: unknown) =>
|
|
5
|
+
new Request("https://example.test/api/adk", {
|
|
6
|
+
method: "POST",
|
|
7
|
+
body: JSON.stringify(body),
|
|
8
|
+
});
|
|
9
|
+
|
|
10
|
+
describe("createAdkApiRoute request controls", () => {
|
|
11
|
+
it("does not trust client run configuration or state by default", async () => {
|
|
12
|
+
const runAsync = vi.fn(async function* () {});
|
|
13
|
+
const handler = createAdkApiRoute({
|
|
14
|
+
runner: { runAsync },
|
|
15
|
+
userId: "user-1",
|
|
16
|
+
sessionId: "session-1",
|
|
17
|
+
});
|
|
18
|
+
|
|
19
|
+
await handler(
|
|
20
|
+
makeRequest({
|
|
21
|
+
message: "Hello",
|
|
22
|
+
runConfig: { maxLlmCalls: 1_000_000 },
|
|
23
|
+
stateDelta: { "app:role": "admin" },
|
|
24
|
+
}),
|
|
25
|
+
);
|
|
26
|
+
|
|
27
|
+
expect(runAsync).toHaveBeenCalledWith({
|
|
28
|
+
userId: "user-1",
|
|
29
|
+
sessionId: "session-1",
|
|
30
|
+
newMessage: { role: "user", parts: [{ text: "Hello" }] },
|
|
31
|
+
});
|
|
32
|
+
});
|
|
33
|
+
|
|
34
|
+
it("forwards only values returned by server resolvers", async () => {
|
|
35
|
+
const runAsync = vi.fn(async function* () {});
|
|
36
|
+
const resolveRunConfig = vi.fn(() => ({ maxLlmCalls: 10 }));
|
|
37
|
+
const resolveStateDelta = vi.fn(() => ({ taskId: "server-task" }));
|
|
38
|
+
const handler = createAdkApiRoute({
|
|
39
|
+
runner: { runAsync },
|
|
40
|
+
userId: "user-1",
|
|
41
|
+
sessionId: "session-1",
|
|
42
|
+
resolveRunConfig,
|
|
43
|
+
resolveStateDelta,
|
|
44
|
+
});
|
|
45
|
+
const request = makeRequest({
|
|
46
|
+
message: "Hello",
|
|
47
|
+
runConfig: { maxLlmCalls: 1_000_000 },
|
|
48
|
+
stateDelta: { "app:role": "admin" },
|
|
49
|
+
});
|
|
50
|
+
|
|
51
|
+
await handler(request);
|
|
52
|
+
|
|
53
|
+
expect(resolveRunConfig).toHaveBeenCalledWith(request, {
|
|
54
|
+
maxLlmCalls: 1_000_000,
|
|
55
|
+
});
|
|
56
|
+
expect(resolveStateDelta).toHaveBeenCalledWith(request, {
|
|
57
|
+
"app:role": "admin",
|
|
58
|
+
});
|
|
59
|
+
expect(runAsync).toHaveBeenCalledWith(
|
|
60
|
+
expect.objectContaining({
|
|
61
|
+
runConfig: { maxLlmCalls: 10 },
|
|
62
|
+
stateDelta: { taskId: "server-task" },
|
|
63
|
+
}),
|
|
64
|
+
);
|
|
65
|
+
});
|
|
66
|
+
});
|
|
@@ -0,0 +1,282 @@
|
|
|
1
|
+
import { describe, expect, it, vi } from "vitest";
|
|
2
|
+
import { createAdkApiRoute } from "./createAdkApiRoute";
|
|
3
|
+
|
|
4
|
+
describe("createAdkApiRoute", () => {
|
|
5
|
+
it("creates the client thread session before the first run", async () => {
|
|
6
|
+
const sessions = new Set<string>();
|
|
7
|
+
const getSession = vi.fn(async ({ sessionId }: { sessionId: string }) =>
|
|
8
|
+
sessions.has(sessionId) ? { id: sessionId } : undefined,
|
|
9
|
+
);
|
|
10
|
+
const createSession = vi.fn(
|
|
11
|
+
async ({ sessionId }: { sessionId: string }) => {
|
|
12
|
+
sessions.add(sessionId);
|
|
13
|
+
return { id: sessionId };
|
|
14
|
+
},
|
|
15
|
+
);
|
|
16
|
+
const runAsync = vi.fn((options: Record<string, unknown>) => {
|
|
17
|
+
expect(sessions.has(options.sessionId as string)).toBe(true);
|
|
18
|
+
return (async function* () {})();
|
|
19
|
+
});
|
|
20
|
+
const sessionId = vi.fn(
|
|
21
|
+
(_request: Request, clientSessionId: string | undefined) => {
|
|
22
|
+
if (!clientSessionId) throw new Error("Missing session ID");
|
|
23
|
+
return `scoped-${clientSessionId}`;
|
|
24
|
+
},
|
|
25
|
+
);
|
|
26
|
+
const handler = createAdkApiRoute({
|
|
27
|
+
runner: {
|
|
28
|
+
appName: "test-app",
|
|
29
|
+
sessionService: { getSession, createSession },
|
|
30
|
+
runAsync,
|
|
31
|
+
},
|
|
32
|
+
userId: "user-1",
|
|
33
|
+
sessionId,
|
|
34
|
+
});
|
|
35
|
+
|
|
36
|
+
await handler(
|
|
37
|
+
new Request("https://example.test/api/adk", {
|
|
38
|
+
method: "POST",
|
|
39
|
+
body: JSON.stringify({ message: "Hello", sessionId: "thread-1" }),
|
|
40
|
+
}),
|
|
41
|
+
);
|
|
42
|
+
|
|
43
|
+
expect(sessionId).toHaveBeenCalledWith(expect.any(Request), "thread-1");
|
|
44
|
+
expect(createSession).toHaveBeenCalledWith({
|
|
45
|
+
appName: "test-app",
|
|
46
|
+
userId: "user-1",
|
|
47
|
+
sessionId: "scoped-thread-1",
|
|
48
|
+
});
|
|
49
|
+
expect(runAsync).toHaveBeenCalledWith(
|
|
50
|
+
expect.objectContaining({
|
|
51
|
+
userId: "user-1",
|
|
52
|
+
sessionId: "scoped-thread-1",
|
|
53
|
+
}),
|
|
54
|
+
);
|
|
55
|
+
});
|
|
56
|
+
|
|
57
|
+
it("uses a separate session for each thread of the same user", async () => {
|
|
58
|
+
const getSession = vi.fn(async () => undefined);
|
|
59
|
+
const createSession = vi.fn(
|
|
60
|
+
async ({ sessionId }: { sessionId: string }) => ({
|
|
61
|
+
id: sessionId,
|
|
62
|
+
}),
|
|
63
|
+
);
|
|
64
|
+
const runAsync = vi.fn(async function* (
|
|
65
|
+
_options: Record<string, unknown>,
|
|
66
|
+
) {});
|
|
67
|
+
const handler = createAdkApiRoute({
|
|
68
|
+
runner: {
|
|
69
|
+
appName: "test-app",
|
|
70
|
+
sessionService: { getSession, createSession },
|
|
71
|
+
runAsync,
|
|
72
|
+
},
|
|
73
|
+
userId: "user-1",
|
|
74
|
+
sessionId: (_request, clientSessionId) => {
|
|
75
|
+
if (!clientSessionId) throw new Error("Missing session ID");
|
|
76
|
+
return clientSessionId;
|
|
77
|
+
},
|
|
78
|
+
});
|
|
79
|
+
|
|
80
|
+
for (const threadId of ["thread-1", "thread-2"]) {
|
|
81
|
+
await handler(
|
|
82
|
+
new Request("https://example.test/api/adk", {
|
|
83
|
+
method: "POST",
|
|
84
|
+
body: JSON.stringify({ message: "Hello", sessionId: threadId }),
|
|
85
|
+
}),
|
|
86
|
+
);
|
|
87
|
+
}
|
|
88
|
+
|
|
89
|
+
expect(getSession).toHaveBeenCalledTimes(2);
|
|
90
|
+
expect(getSession).toHaveBeenNthCalledWith(1, {
|
|
91
|
+
appName: "test-app",
|
|
92
|
+
userId: "user-1",
|
|
93
|
+
sessionId: "thread-1",
|
|
94
|
+
});
|
|
95
|
+
expect(getSession).toHaveBeenNthCalledWith(2, {
|
|
96
|
+
appName: "test-app",
|
|
97
|
+
userId: "user-1",
|
|
98
|
+
sessionId: "thread-2",
|
|
99
|
+
});
|
|
100
|
+
expect(createSession).toHaveBeenCalledTimes(2);
|
|
101
|
+
expect(createSession).toHaveBeenNthCalledWith(1, {
|
|
102
|
+
appName: "test-app",
|
|
103
|
+
userId: "user-1",
|
|
104
|
+
sessionId: "thread-1",
|
|
105
|
+
});
|
|
106
|
+
expect(createSession).toHaveBeenNthCalledWith(2, {
|
|
107
|
+
appName: "test-app",
|
|
108
|
+
userId: "user-1",
|
|
109
|
+
sessionId: "thread-2",
|
|
110
|
+
});
|
|
111
|
+
expect(runAsync).toHaveBeenCalledTimes(2);
|
|
112
|
+
expect(runAsync.mock.calls[0]?.[0]).toMatchObject({
|
|
113
|
+
userId: "user-1",
|
|
114
|
+
sessionId: "thread-1",
|
|
115
|
+
});
|
|
116
|
+
expect(runAsync.mock.calls[1]?.[0]).toMatchObject({
|
|
117
|
+
userId: "user-1",
|
|
118
|
+
sessionId: "thread-2",
|
|
119
|
+
});
|
|
120
|
+
});
|
|
121
|
+
|
|
122
|
+
it("resolves a supplied session ID only within the authenticated user", async () => {
|
|
123
|
+
const getSession = vi.fn(async ({ userId }: { userId: string }) =>
|
|
124
|
+
userId === "user-a" ? { id: "user-a-session" } : undefined,
|
|
125
|
+
);
|
|
126
|
+
const createSession = vi.fn(
|
|
127
|
+
async ({ sessionId }: { sessionId: string }) => ({
|
|
128
|
+
id: sessionId,
|
|
129
|
+
}),
|
|
130
|
+
);
|
|
131
|
+
const runAsync = vi.fn(async function* () {});
|
|
132
|
+
const authenticatedUsers = new WeakMap<Request, string>();
|
|
133
|
+
const handler = createAdkApiRoute({
|
|
134
|
+
runner: {
|
|
135
|
+
appName: "test-app",
|
|
136
|
+
sessionService: { getSession, createSession },
|
|
137
|
+
runAsync,
|
|
138
|
+
},
|
|
139
|
+
userId: (request) => {
|
|
140
|
+
const userId = authenticatedUsers.get(request);
|
|
141
|
+
if (!userId) throw new Error("Unauthenticated");
|
|
142
|
+
return userId;
|
|
143
|
+
},
|
|
144
|
+
sessionId: (_request, clientSessionId) => {
|
|
145
|
+
if (!clientSessionId) throw new Error("Missing session ID");
|
|
146
|
+
return clientSessionId;
|
|
147
|
+
},
|
|
148
|
+
});
|
|
149
|
+
const request = new Request("https://example.test/api/adk", {
|
|
150
|
+
method: "POST",
|
|
151
|
+
body: JSON.stringify({ message: "Hello", sessionId: "user-a-session" }),
|
|
152
|
+
});
|
|
153
|
+
authenticatedUsers.set(request, "user-b");
|
|
154
|
+
|
|
155
|
+
await handler(request);
|
|
156
|
+
|
|
157
|
+
expect(getSession).toHaveBeenCalledExactlyOnceWith({
|
|
158
|
+
appName: "test-app",
|
|
159
|
+
userId: "user-b",
|
|
160
|
+
sessionId: "user-a-session",
|
|
161
|
+
});
|
|
162
|
+
expect(createSession).toHaveBeenCalledExactlyOnceWith({
|
|
163
|
+
appName: "test-app",
|
|
164
|
+
userId: "user-b",
|
|
165
|
+
sessionId: "user-a-session",
|
|
166
|
+
});
|
|
167
|
+
expect(runAsync).toHaveBeenCalledWith(
|
|
168
|
+
expect.objectContaining({
|
|
169
|
+
userId: "user-b",
|
|
170
|
+
sessionId: "user-a-session",
|
|
171
|
+
}),
|
|
172
|
+
);
|
|
173
|
+
});
|
|
174
|
+
|
|
175
|
+
it("shares session creation across concurrent first requests", async () => {
|
|
176
|
+
const getSession = vi.fn(async () => undefined);
|
|
177
|
+
let releaseCreation!: () => void;
|
|
178
|
+
let markCreationStarted!: () => void;
|
|
179
|
+
const creationStarted = new Promise<void>((resolve) => {
|
|
180
|
+
markCreationStarted = resolve;
|
|
181
|
+
});
|
|
182
|
+
const createSession = vi.fn(async () => {
|
|
183
|
+
markCreationStarted();
|
|
184
|
+
await new Promise<void>((resolve) => {
|
|
185
|
+
releaseCreation = resolve;
|
|
186
|
+
});
|
|
187
|
+
return { id: "thread-1" };
|
|
188
|
+
});
|
|
189
|
+
const runAsync = vi.fn(async function* () {});
|
|
190
|
+
const sessionId = vi.fn(
|
|
191
|
+
(_request: Request, clientSessionId: string | undefined) =>
|
|
192
|
+
clientSessionId ?? "missing",
|
|
193
|
+
);
|
|
194
|
+
const handler = createAdkApiRoute({
|
|
195
|
+
runner: {
|
|
196
|
+
appName: "test-app",
|
|
197
|
+
sessionService: { getSession, createSession },
|
|
198
|
+
runAsync,
|
|
199
|
+
},
|
|
200
|
+
userId: "user-1",
|
|
201
|
+
sessionId,
|
|
202
|
+
});
|
|
203
|
+
const makeRequest = () =>
|
|
204
|
+
new Request("https://example.test/api/adk", {
|
|
205
|
+
method: "POST",
|
|
206
|
+
body: JSON.stringify({ message: "Hello", sessionId: "thread-1" }),
|
|
207
|
+
});
|
|
208
|
+
|
|
209
|
+
const firstRequest = handler(makeRequest());
|
|
210
|
+
await creationStarted;
|
|
211
|
+
const secondRequest = handler(makeRequest());
|
|
212
|
+
await vi.waitFor(() => expect(sessionId).toHaveBeenCalledTimes(2));
|
|
213
|
+
await Promise.resolve();
|
|
214
|
+
releaseCreation();
|
|
215
|
+
await Promise.all([firstRequest, secondRequest]);
|
|
216
|
+
|
|
217
|
+
expect(getSession).toHaveBeenCalledTimes(1);
|
|
218
|
+
expect(createSession).toHaveBeenCalledTimes(1);
|
|
219
|
+
expect(runAsync).toHaveBeenCalledTimes(2);
|
|
220
|
+
});
|
|
221
|
+
|
|
222
|
+
it("continues when another process creates the session first", async () => {
|
|
223
|
+
let exists = false;
|
|
224
|
+
const getSession = vi.fn(async () =>
|
|
225
|
+
exists ? { id: "thread-1" } : undefined,
|
|
226
|
+
);
|
|
227
|
+
const createSession = vi.fn(async () => {
|
|
228
|
+
exists = true;
|
|
229
|
+
throw new Error("Session already exists");
|
|
230
|
+
});
|
|
231
|
+
const runAsync = vi.fn(async function* () {});
|
|
232
|
+
const handler = createAdkApiRoute({
|
|
233
|
+
runner: {
|
|
234
|
+
appName: "test-app",
|
|
235
|
+
sessionService: { getSession, createSession },
|
|
236
|
+
runAsync,
|
|
237
|
+
},
|
|
238
|
+
userId: "user-1",
|
|
239
|
+
sessionId: "thread-1",
|
|
240
|
+
});
|
|
241
|
+
|
|
242
|
+
await handler(
|
|
243
|
+
new Request("https://example.test/api/adk", {
|
|
244
|
+
method: "POST",
|
|
245
|
+
body: JSON.stringify({ message: "Hello" }),
|
|
246
|
+
}),
|
|
247
|
+
);
|
|
248
|
+
|
|
249
|
+
expect(getSession).toHaveBeenCalledTimes(2);
|
|
250
|
+
expect(createSession).toHaveBeenCalledOnce();
|
|
251
|
+
expect(runAsync).toHaveBeenCalledOnce();
|
|
252
|
+
});
|
|
253
|
+
|
|
254
|
+
it("preserves create failures when the session still does not exist", async () => {
|
|
255
|
+
const failure = new Error("database unavailable");
|
|
256
|
+
const getSession = vi.fn(async () => undefined);
|
|
257
|
+
const createSession = vi.fn(async () => {
|
|
258
|
+
throw failure;
|
|
259
|
+
});
|
|
260
|
+
const runAsync = vi.fn(async function* () {});
|
|
261
|
+
const handler = createAdkApiRoute({
|
|
262
|
+
runner: {
|
|
263
|
+
appName: "test-app",
|
|
264
|
+
sessionService: { getSession, createSession },
|
|
265
|
+
runAsync,
|
|
266
|
+
},
|
|
267
|
+
userId: "user-1",
|
|
268
|
+
sessionId: "thread-1",
|
|
269
|
+
});
|
|
270
|
+
|
|
271
|
+
await expect(
|
|
272
|
+
handler(
|
|
273
|
+
new Request("https://example.test/api/adk", {
|
|
274
|
+
method: "POST",
|
|
275
|
+
body: JSON.stringify({ message: "Hello" }),
|
|
276
|
+
}),
|
|
277
|
+
),
|
|
278
|
+
).rejects.toBe(failure);
|
|
279
|
+
expect(getSession).toHaveBeenCalledTimes(2);
|
|
280
|
+
expect(runAsync).not.toHaveBeenCalled();
|
|
281
|
+
});
|
|
282
|
+
});
|
|
@@ -6,11 +6,77 @@ import { adkEventStream, type AdkEventStreamOptions } from "./adkEventStream";
|
|
|
6
6
|
* Avoids requiring `@google/adk` as a dependency.
|
|
7
7
|
*/
|
|
8
8
|
type AdkRunner = {
|
|
9
|
+
readonly appName?: string;
|
|
10
|
+
readonly sessionService?: {
|
|
11
|
+
getSession(options: {
|
|
12
|
+
appName: string;
|
|
13
|
+
userId: string;
|
|
14
|
+
sessionId: string;
|
|
15
|
+
}): Promise<unknown | undefined>;
|
|
16
|
+
createSession(options: {
|
|
17
|
+
appName: string;
|
|
18
|
+
userId: string;
|
|
19
|
+
sessionId: string;
|
|
20
|
+
}): Promise<unknown>;
|
|
21
|
+
};
|
|
9
22
|
runAsync(
|
|
10
23
|
options: Record<string, unknown>,
|
|
11
24
|
): AsyncGenerator<any, void, undefined>;
|
|
12
25
|
};
|
|
13
26
|
|
|
27
|
+
type AdkSessionService = NonNullable<AdkRunner["sessionService"]>;
|
|
28
|
+
|
|
29
|
+
const pendingSessions = new WeakMap<
|
|
30
|
+
AdkSessionService,
|
|
31
|
+
Map<string, Promise<void>>
|
|
32
|
+
>();
|
|
33
|
+
|
|
34
|
+
const ensureRunnerSession = async (
|
|
35
|
+
runner: AdkRunner,
|
|
36
|
+
userId: string,
|
|
37
|
+
sessionId: string,
|
|
38
|
+
) => {
|
|
39
|
+
const { appName, sessionService } = runner;
|
|
40
|
+
if (!appName || !sessionService) return;
|
|
41
|
+
|
|
42
|
+
let serviceSessions = pendingSessions.get(sessionService);
|
|
43
|
+
if (!serviceSessions) {
|
|
44
|
+
serviceSessions = new Map();
|
|
45
|
+
pendingSessions.set(sessionService, serviceSessions);
|
|
46
|
+
}
|
|
47
|
+
|
|
48
|
+
const key = JSON.stringify([appName, userId, sessionId]);
|
|
49
|
+
let pending = serviceSessions.get(key);
|
|
50
|
+
if (!pending) {
|
|
51
|
+
pending = (async () => {
|
|
52
|
+
const session = await sessionService.getSession({
|
|
53
|
+
appName,
|
|
54
|
+
userId,
|
|
55
|
+
sessionId,
|
|
56
|
+
});
|
|
57
|
+
if (!session) {
|
|
58
|
+
try {
|
|
59
|
+
await sessionService.createSession({ appName, userId, sessionId });
|
|
60
|
+
} catch (error) {
|
|
61
|
+
const existing = await sessionService.getSession({
|
|
62
|
+
appName,
|
|
63
|
+
userId,
|
|
64
|
+
sessionId,
|
|
65
|
+
});
|
|
66
|
+
if (!existing) throw error;
|
|
67
|
+
}
|
|
68
|
+
}
|
|
69
|
+
})();
|
|
70
|
+
serviceSessions.set(key, pending);
|
|
71
|
+
}
|
|
72
|
+
|
|
73
|
+
try {
|
|
74
|
+
await pending;
|
|
75
|
+
} finally {
|
|
76
|
+
if (serviceSessions.get(key) === pending) serviceSessions.delete(key);
|
|
77
|
+
}
|
|
78
|
+
};
|
|
79
|
+
|
|
14
80
|
export type CreateAdkApiRouteOptions = {
|
|
15
81
|
/**
|
|
16
82
|
* ADK Runner instance.
|
|
@@ -18,16 +84,47 @@ export type CreateAdkApiRouteOptions = {
|
|
|
18
84
|
runner: AdkRunner;
|
|
19
85
|
|
|
20
86
|
/**
|
|
21
|
-
* User ID to use for the ADK session.
|
|
22
|
-
*
|
|
87
|
+
* User ID to use for the ADK session. Production routes should resolve this
|
|
88
|
+
* from the authenticated request. A static value is suitable only for a
|
|
89
|
+
* single-user development route.
|
|
23
90
|
*/
|
|
24
91
|
userId: string | ((req: Request) => string | Promise<string>);
|
|
25
92
|
|
|
26
93
|
/**
|
|
27
94
|
* Session ID to use. Can be a static string or a function
|
|
28
|
-
* that
|
|
95
|
+
* that validates or transforms the client thread ID sent by
|
|
96
|
+
* `createAdkStream`. The client value is an identifier, not authorization;
|
|
97
|
+
* scope access with an authenticated `userId`.
|
|
29
98
|
*/
|
|
30
|
-
sessionId:
|
|
99
|
+
sessionId:
|
|
100
|
+
| string
|
|
101
|
+
| ((
|
|
102
|
+
req: Request,
|
|
103
|
+
clientSessionId: string | undefined,
|
|
104
|
+
) => string | Promise<string>);
|
|
105
|
+
|
|
106
|
+
/**
|
|
107
|
+
* Validates or replaces the client-provided ADK run configuration.
|
|
108
|
+
* Client values are ignored unless this resolver is provided.
|
|
109
|
+
*/
|
|
110
|
+
resolveRunConfig?:
|
|
111
|
+
| ((req: Request, runConfig: unknown) => unknown | Promise<unknown>)
|
|
112
|
+
| undefined;
|
|
113
|
+
|
|
114
|
+
/**
|
|
115
|
+
* Validates or replaces the client-provided ADK state delta.
|
|
116
|
+
* Client values are ignored unless this resolver is provided. In particular,
|
|
117
|
+
* `app:` and `user:` keys affect state beyond the current session.
|
|
118
|
+
*/
|
|
119
|
+
resolveStateDelta?:
|
|
120
|
+
| ((
|
|
121
|
+
req: Request,
|
|
122
|
+
stateDelta: Record<string, unknown> | undefined,
|
|
123
|
+
) =>
|
|
124
|
+
| Record<string, unknown>
|
|
125
|
+
| undefined
|
|
126
|
+
| Promise<Record<string, unknown> | undefined>)
|
|
127
|
+
| undefined;
|
|
31
128
|
|
|
32
129
|
/**
|
|
33
130
|
* Error handler for stream errors.
|
|
@@ -43,11 +140,15 @@ export type CreateAdkApiRouteOptions = {
|
|
|
43
140
|
* ```ts
|
|
44
141
|
* import { createAdkApiRoute } from '@assistant-ui/react-google-adk/server';
|
|
45
142
|
* import { runner } from './agent';
|
|
143
|
+
* import { requireUser } from './auth';
|
|
46
144
|
*
|
|
47
145
|
* export const POST = createAdkApiRoute({
|
|
48
146
|
* runner,
|
|
49
|
-
* userId:
|
|
50
|
-
* sessionId: (
|
|
147
|
+
* userId: async (req) => (await requireUser(req)).id,
|
|
148
|
+
* sessionId: (_req, clientSessionId) => {
|
|
149
|
+
* if (!clientSessionId) throw new Error("Missing ADK session ID");
|
|
150
|
+
* return clientSessionId;
|
|
151
|
+
* },
|
|
51
152
|
* });
|
|
52
153
|
* ```
|
|
53
154
|
*/
|
|
@@ -65,17 +166,24 @@ export function createAdkApiRoute(
|
|
|
65
166
|
|
|
66
167
|
const sessionId =
|
|
67
168
|
typeof options.sessionId === "function"
|
|
68
|
-
? await options.sessionId(req)
|
|
169
|
+
? await options.sessionId(req, parsed.sessionId)
|
|
69
170
|
: options.sessionId;
|
|
70
171
|
|
|
172
|
+
const runConfig = options.resolveRunConfig
|
|
173
|
+
? await options.resolveRunConfig(req, parsed.config.runConfig)
|
|
174
|
+
: undefined;
|
|
175
|
+
const stateDelta = options.resolveStateDelta
|
|
176
|
+
? await options.resolveStateDelta(req, parsed.stateDelta)
|
|
177
|
+
: undefined;
|
|
178
|
+
|
|
179
|
+
await ensureRunnerSession(options.runner, userId, sessionId);
|
|
180
|
+
|
|
71
181
|
const events = options.runner.runAsync({
|
|
72
182
|
userId,
|
|
73
183
|
sessionId,
|
|
74
184
|
newMessage,
|
|
75
|
-
...(
|
|
76
|
-
...(
|
|
77
|
-
runConfig: parsed.config.runConfig,
|
|
78
|
-
}),
|
|
185
|
+
...(stateDelta != null && { stateDelta }),
|
|
186
|
+
...(runConfig != null && { runConfig }),
|
|
79
187
|
});
|
|
80
188
|
|
|
81
189
|
return adkEventStream(
|
|
@@ -28,11 +28,18 @@ describe("parseAdkRequest", () => {
|
|
|
28
28
|
});
|
|
29
29
|
});
|
|
30
30
|
|
|
31
|
-
it("parses a message request with stateDelta", async () => {
|
|
31
|
+
it("parses a message request with sessionId and stateDelta", async () => {
|
|
32
32
|
const result = await parseAdkRequest(
|
|
33
|
-
makeRequest({
|
|
33
|
+
makeRequest({
|
|
34
|
+
message: "Hello",
|
|
35
|
+
sessionId: "session-1",
|
|
36
|
+
stateDelta: { count: 1 },
|
|
37
|
+
}),
|
|
34
38
|
);
|
|
35
|
-
expect(result).toMatchObject({
|
|
39
|
+
expect(result).toMatchObject({
|
|
40
|
+
sessionId: "session-1",
|
|
41
|
+
stateDelta: { count: 1 },
|
|
42
|
+
});
|
|
36
43
|
});
|
|
37
44
|
|
|
38
45
|
it("parses a message request with parts (multimodal)", async () => {
|
|
@@ -130,6 +137,7 @@ describe("parseAdkRequest", () => {
|
|
|
130
137
|
[{ parts: [null] }, 'field "parts"'],
|
|
131
138
|
[{ type: "unknown", message: "hello" }, 'field "type"'],
|
|
132
139
|
[{ message: "hello", checkpointId: 42 }, 'field "checkpointId"'],
|
|
140
|
+
[{ message: "hello", sessionId: 42 }, 'field "sessionId"'],
|
|
133
141
|
[{ message: "hello", stateDelta: [] }, 'field "stateDelta"'],
|
|
134
142
|
])("rejects malformed message requests %#", async (body, error) => {
|
|
135
143
|
await expect(parseAdkRequest(makeRequest(body))).rejects.toThrow(error);
|
|
@@ -7,6 +7,7 @@ type ParsedAdkRequest =
|
|
|
7
7
|
type: "message";
|
|
8
8
|
text: string;
|
|
9
9
|
parts?: Array<Record<string, unknown>> | undefined;
|
|
10
|
+
sessionId?: string | undefined;
|
|
10
11
|
config: AdkSendMessageConfig;
|
|
11
12
|
stateDelta?: Record<string, unknown> | undefined;
|
|
12
13
|
}
|
|
@@ -16,6 +17,7 @@ type ParsedAdkRequest =
|
|
|
16
17
|
toolName: string;
|
|
17
18
|
result: unknown;
|
|
18
19
|
isError: boolean;
|
|
20
|
+
sessionId?: string | undefined;
|
|
19
21
|
config: AdkSendMessageConfig;
|
|
20
22
|
stateDelta?: Record<string, unknown> | undefined;
|
|
21
23
|
};
|
|
@@ -162,6 +164,7 @@ export const parseAdkRequest = async (
|
|
|
162
164
|
if (body.runConfig !== undefined) config.runConfig = body.runConfig;
|
|
163
165
|
const checkpointId = readOptionalString(body, "checkpointId");
|
|
164
166
|
if (checkpointId !== undefined) config.checkpointId = checkpointId;
|
|
167
|
+
const sessionId = readOptionalString(body, "sessionId");
|
|
165
168
|
|
|
166
169
|
const stateDelta = body.stateDelta;
|
|
167
170
|
if (stateDelta !== undefined && !isRecord(stateDelta)) {
|
|
@@ -181,6 +184,7 @@ export const parseAdkRequest = async (
|
|
|
181
184
|
toolName: readString(body, "toolName"),
|
|
182
185
|
result: body.result,
|
|
183
186
|
isError: body.isError ?? false,
|
|
187
|
+
...(sessionId !== undefined && { sessionId }),
|
|
184
188
|
config,
|
|
185
189
|
...(stateDelta != null && { stateDelta }),
|
|
186
190
|
};
|
|
@@ -213,6 +217,7 @@ export const parseAdkRequest = async (
|
|
|
213
217
|
type: "message",
|
|
214
218
|
text: text ?? "",
|
|
215
219
|
...(parts !== undefined && { parts }),
|
|
220
|
+
...(sessionId !== undefined && { sessionId }),
|
|
216
221
|
config,
|
|
217
222
|
...(stateDelta != null && { stateDelta }),
|
|
218
223
|
};
|
|
@@ -226,7 +231,8 @@ export const parseAdkRequest = async (
|
|
|
226
231
|
* ```ts
|
|
227
232
|
* const parsed = await parseAdkRequest(req);
|
|
228
233
|
* const newMessage = toAdkContent(parsed);
|
|
229
|
-
* const
|
|
234
|
+
* const stateDelta = validateSessionState(parsed.stateDelta);
|
|
235
|
+
* const events = runner.runAsync({ userId, sessionId, newMessage, stateDelta });
|
|
230
236
|
* return adkEventStream(events);
|
|
231
237
|
* ```
|
|
232
238
|
*/
|