@assistant-ui/react-google-adk 0.0.34 → 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 +29 -4
- package/dist/useAdkMessages.d.ts.map +1 -1
- package/dist/useAdkMessages.js +62 -76
- 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 +182 -54
- package/dist/useAdkRuntime.js.map +1 -1
- package/package.json +6 -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 +43 -0
- package/src/useAdkMessages.ts +89 -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 +253 -75
- package/src/useAdkRuntimeApproval.test.tsx +390 -35
- 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
|
@@ -0,0 +1,124 @@
|
|
|
1
|
+
import type { AppendMessage } from "@assistant-ui/core";
|
|
2
|
+
import type {
|
|
3
|
+
AdkAuthRequest,
|
|
4
|
+
AdkMessage,
|
|
5
|
+
AdkMessageMetadata,
|
|
6
|
+
AdkThreadSnapshot,
|
|
7
|
+
AdkToolConfirmation,
|
|
8
|
+
} from "./types";
|
|
9
|
+
|
|
10
|
+
export type AdkStagedEntry = {
|
|
11
|
+
message: AdkMessage & { id: string };
|
|
12
|
+
runConfig: AppendMessage["runConfig"];
|
|
13
|
+
};
|
|
14
|
+
|
|
15
|
+
export type AdkThreadState = {
|
|
16
|
+
messages: AdkMessage[];
|
|
17
|
+
stateDelta: Record<string, unknown>;
|
|
18
|
+
agentInfo: { name?: string | undefined; branch?: string | undefined };
|
|
19
|
+
longRunningToolIds: string[];
|
|
20
|
+
artifactDelta: Record<string, number>;
|
|
21
|
+
toolConfirmations: AdkToolConfirmation[];
|
|
22
|
+
authRequests: AdkAuthRequest[];
|
|
23
|
+
escalated: boolean;
|
|
24
|
+
messageMetadata: Map<string, AdkMessageMetadata>;
|
|
25
|
+
stagedEntries: ReadonlyMap<string, AdkStagedEntry>;
|
|
26
|
+
};
|
|
27
|
+
|
|
28
|
+
export type AdkThreadAction =
|
|
29
|
+
| { type: "event.published"; state: Omit<AdkThreadState, "stagedEntries"> }
|
|
30
|
+
| { type: "snapshot.applied"; snapshot: AdkThreadSnapshot }
|
|
31
|
+
| { type: "messages.replaced"; messages: AdkMessage[] }
|
|
32
|
+
| { type: "messages.set"; messages: AdkMessage[] }
|
|
33
|
+
| { type: "longRunningToolIds.set"; ids: string[] }
|
|
34
|
+
| { type: "staged.stage"; entry: AdkStagedEntry }
|
|
35
|
+
| { type: "staged.unstage"; ids: readonly string[] }
|
|
36
|
+
| {
|
|
37
|
+
type: "run.started";
|
|
38
|
+
messages: AdkMessage[];
|
|
39
|
+
longRunningToolIds: string[];
|
|
40
|
+
toolConfirmations: AdkToolConfirmation[];
|
|
41
|
+
authRequests: AdkAuthRequest[];
|
|
42
|
+
};
|
|
43
|
+
|
|
44
|
+
export const createAdkThreadState = (): AdkThreadState => ({
|
|
45
|
+
messages: [],
|
|
46
|
+
stateDelta: {},
|
|
47
|
+
agentInfo: {},
|
|
48
|
+
longRunningToolIds: [],
|
|
49
|
+
artifactDelta: {},
|
|
50
|
+
toolConfirmations: [],
|
|
51
|
+
authRequests: [],
|
|
52
|
+
escalated: false,
|
|
53
|
+
messageMetadata: new Map(),
|
|
54
|
+
stagedEntries: new Map(),
|
|
55
|
+
});
|
|
56
|
+
|
|
57
|
+
export const reduceAdkThreadState = (
|
|
58
|
+
state: AdkThreadState,
|
|
59
|
+
action: AdkThreadAction,
|
|
60
|
+
): AdkThreadState => {
|
|
61
|
+
switch (action.type) {
|
|
62
|
+
case "event.published": {
|
|
63
|
+
const next = action.state;
|
|
64
|
+
return {
|
|
65
|
+
...state,
|
|
66
|
+
...next,
|
|
67
|
+
stateDelta: { ...state.stateDelta, ...next.stateDelta },
|
|
68
|
+
artifactDelta: { ...state.artifactDelta, ...next.artifactDelta },
|
|
69
|
+
messageMetadata:
|
|
70
|
+
next.messageMetadata.size > 0
|
|
71
|
+
? new Map([...state.messageMetadata, ...next.messageMetadata])
|
|
72
|
+
: state.messageMetadata,
|
|
73
|
+
};
|
|
74
|
+
}
|
|
75
|
+
case "snapshot.applied": {
|
|
76
|
+
const snapshot = action.snapshot;
|
|
77
|
+
return {
|
|
78
|
+
...state,
|
|
79
|
+
messages: snapshot.messages,
|
|
80
|
+
stateDelta: snapshot.stateDelta ?? {},
|
|
81
|
+
agentInfo: snapshot.agentInfo ?? {},
|
|
82
|
+
longRunningToolIds: snapshot.longRunningToolIds ?? [],
|
|
83
|
+
artifactDelta: snapshot.artifactDelta ?? {},
|
|
84
|
+
toolConfirmations: snapshot.toolConfirmations ?? [],
|
|
85
|
+
authRequests: snapshot.authRequests ?? [],
|
|
86
|
+
escalated: snapshot.escalated ?? false,
|
|
87
|
+
messageMetadata: snapshot.messageMetadata ?? new Map(),
|
|
88
|
+
};
|
|
89
|
+
}
|
|
90
|
+
case "messages.replaced":
|
|
91
|
+
return {
|
|
92
|
+
...state,
|
|
93
|
+
messages: action.messages,
|
|
94
|
+
longRunningToolIds: [],
|
|
95
|
+
toolConfirmations: [],
|
|
96
|
+
authRequests: [],
|
|
97
|
+
escalated: false,
|
|
98
|
+
messageMetadata: new Map(),
|
|
99
|
+
};
|
|
100
|
+
case "messages.set":
|
|
101
|
+
return { ...state, messages: action.messages };
|
|
102
|
+
case "longRunningToolIds.set":
|
|
103
|
+
return { ...state, longRunningToolIds: action.ids };
|
|
104
|
+
case "staged.stage": {
|
|
105
|
+
const stagedEntries = new Map(state.stagedEntries);
|
|
106
|
+
stagedEntries.set(action.entry.message.id, action.entry);
|
|
107
|
+
return { ...state, stagedEntries };
|
|
108
|
+
}
|
|
109
|
+
case "staged.unstage": {
|
|
110
|
+
if (!action.ids.some((id) => state.stagedEntries.has(id))) return state;
|
|
111
|
+
const stagedEntries = new Map(state.stagedEntries);
|
|
112
|
+
for (const id of action.ids) stagedEntries.delete(id);
|
|
113
|
+
return { ...state, stagedEntries };
|
|
114
|
+
}
|
|
115
|
+
case "run.started":
|
|
116
|
+
return {
|
|
117
|
+
...state,
|
|
118
|
+
messages: action.messages,
|
|
119
|
+
longRunningToolIds: action.longRunningToolIds,
|
|
120
|
+
toolConfirmations: action.toolConfirmations,
|
|
121
|
+
authRequests: action.authRequests,
|
|
122
|
+
};
|
|
123
|
+
}
|
|
124
|
+
};
|
|
@@ -300,6 +300,25 @@ describe("getMessageContent", () => {
|
|
|
300
300
|
]);
|
|
301
301
|
});
|
|
302
302
|
|
|
303
|
+
it("uses the binary data URL default for a media-less file", () => {
|
|
304
|
+
const result = getMessageContent(
|
|
305
|
+
makeAppendMessage([
|
|
306
|
+
{
|
|
307
|
+
type: "file",
|
|
308
|
+
mimeType: "",
|
|
309
|
+
data: "data:;base64,SGVsbG8=",
|
|
310
|
+
},
|
|
311
|
+
]),
|
|
312
|
+
);
|
|
313
|
+
expect(result).toEqual([
|
|
314
|
+
{
|
|
315
|
+
type: "file",
|
|
316
|
+
mimeType: "application/octet-stream",
|
|
317
|
+
data: "SGVsbG8=",
|
|
318
|
+
},
|
|
319
|
+
]);
|
|
320
|
+
});
|
|
321
|
+
|
|
303
322
|
it("emits a file_url part for file parts with sourceType url", () => {
|
|
304
323
|
const result = getMessageContent(
|
|
305
324
|
makeAppendMessage([
|
|
@@ -52,7 +52,7 @@ export const getMessageContent = (msg: AppendMessage) => {
|
|
|
52
52
|
}
|
|
53
53
|
return {
|
|
54
54
|
type: "file" as const,
|
|
55
|
-
mimeType:
|
|
55
|
+
mimeType: source.mimeType,
|
|
56
56
|
// Lands in Gemini `inlineData.data`, which takes bare base64, so a
|
|
57
57
|
// data URL envelope is stripped rather than forwarded.
|
|
58
58
|
data: source.data,
|
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(
|