@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/AdkClient.test.ts
CHANGED
|
@@ -189,7 +189,8 @@ describe("createAdkStream - proxy mode", () => {
|
|
|
189
189
|
const messages: AdkMessage[] = [
|
|
190
190
|
{ id: "m1", type: "human", content: "Hello" },
|
|
191
191
|
];
|
|
192
|
-
const
|
|
192
|
+
const config = makeConfig();
|
|
193
|
+
const gen = await stream(messages, config);
|
|
193
194
|
// drain
|
|
194
195
|
for await (const _ of gen) {
|
|
195
196
|
/* noop */
|
|
@@ -200,7 +201,27 @@ describe("createAdkStream - proxy mode", () => {
|
|
|
200
201
|
expect(url).toBe("/api/adk");
|
|
201
202
|
expect(init?.method).toBe("POST");
|
|
202
203
|
const body = JSON.parse(init?.body as string);
|
|
203
|
-
expect(body).toMatchObject({ message: "Hello" });
|
|
204
|
+
expect(body).toMatchObject({ message: "Hello", sessionId: "session-1" });
|
|
205
|
+
expect(config.initialize).toHaveBeenCalledOnce();
|
|
206
|
+
});
|
|
207
|
+
|
|
208
|
+
it("falls back to the remote thread ID for proxy sessions", async () => {
|
|
209
|
+
mockFetch.mockResolvedValueOnce(sseResponse(sseBody("")));
|
|
210
|
+
const initialize = vi
|
|
211
|
+
.fn()
|
|
212
|
+
.mockResolvedValue({ remoteId: "remote-1", externalId: undefined });
|
|
213
|
+
|
|
214
|
+
const stream = createAdkStream({ api: "/api/adk" });
|
|
215
|
+
const gen = await stream(
|
|
216
|
+
[{ id: "m1", type: "human", content: "Hello" }],
|
|
217
|
+
makeConfig({ initialize }),
|
|
218
|
+
);
|
|
219
|
+
for await (const _ of gen) {
|
|
220
|
+
/* noop */
|
|
221
|
+
}
|
|
222
|
+
|
|
223
|
+
const body = JSON.parse(mockFetch.mock.calls[0]![1]?.body as string);
|
|
224
|
+
expect(body.sessionId).toBe("remote-1");
|
|
204
225
|
});
|
|
205
226
|
|
|
206
227
|
it("sends runConfig and checkpointId in proxy body", async () => {
|
|
@@ -603,6 +624,61 @@ describe("createAdkStream - SSE parsing", () => {
|
|
|
603
624
|
expect(collected[1]!.id).toBe("e2");
|
|
604
625
|
});
|
|
605
626
|
|
|
627
|
+
it("forwards configurable SSE line and event limits", async () => {
|
|
628
|
+
const event: AdkEvent = {
|
|
629
|
+
id: "e1",
|
|
630
|
+
content: { parts: [{ text: "x".repeat(64) }] },
|
|
631
|
+
};
|
|
632
|
+
const text = `data: ${JSON.stringify(event)}\n\n`;
|
|
633
|
+
const consume = async (options: {
|
|
634
|
+
maxStreamLineLength: number;
|
|
635
|
+
maxStreamEventLength: number;
|
|
636
|
+
}) => {
|
|
637
|
+
mockFetch.mockResolvedValueOnce(sseResponse(sseBody(text)));
|
|
638
|
+
const stream = createAdkStream({ api: "/api/adk", ...options });
|
|
639
|
+
const gen = await stream(
|
|
640
|
+
[{ id: "m1", type: "human", content: "Hi" }],
|
|
641
|
+
makeConfig(),
|
|
642
|
+
);
|
|
643
|
+
const events: AdkEvent[] = [];
|
|
644
|
+
for await (const parsedEvent of gen) events.push(parsedEvent);
|
|
645
|
+
return events;
|
|
646
|
+
};
|
|
647
|
+
|
|
648
|
+
await expect(
|
|
649
|
+
consume({ maxStreamLineLength: 32, maxStreamEventLength: 1_024 }),
|
|
650
|
+
).rejects.toThrow("SSE line exceeds maxLineLength");
|
|
651
|
+
await expect(
|
|
652
|
+
consume({ maxStreamLineLength: 1_024, maxStreamEventLength: 32 }),
|
|
653
|
+
).rejects.toThrow("SSE event exceeds maxEventLength");
|
|
654
|
+
await expect(
|
|
655
|
+
consume({ maxStreamLineLength: 1_024, maxStreamEventLength: 1_024 }),
|
|
656
|
+
).resolves.toEqual([event]);
|
|
657
|
+
});
|
|
658
|
+
|
|
659
|
+
it("releases the response reader when decoder limits are invalid", async () => {
|
|
660
|
+
const cancel = vi.fn();
|
|
661
|
+
const body = new ReadableStream<Uint8Array>({ cancel });
|
|
662
|
+
mockFetch.mockResolvedValueOnce(sseResponse(body));
|
|
663
|
+
const stream = createAdkStream({
|
|
664
|
+
api: "/api/adk",
|
|
665
|
+
maxStreamLineLength: 0,
|
|
666
|
+
});
|
|
667
|
+
const consume = async () => {
|
|
668
|
+
const gen = await stream(
|
|
669
|
+
[{ id: "m1", type: "human", content: "Hi" }],
|
|
670
|
+
makeConfig(),
|
|
671
|
+
);
|
|
672
|
+
for await (const _event of gen) void _event;
|
|
673
|
+
};
|
|
674
|
+
|
|
675
|
+
await expect(consume()).rejects.toThrow(
|
|
676
|
+
"maxLineLength must be a positive safe integer",
|
|
677
|
+
);
|
|
678
|
+
expect(cancel).toHaveBeenCalledOnce();
|
|
679
|
+
expect(body.locked).toBe(false);
|
|
680
|
+
});
|
|
681
|
+
|
|
606
682
|
it.each(["{}", "null", "[]", '["event"]', '"event"'])(
|
|
607
683
|
"rejects empty or non-object stream events: %s",
|
|
608
684
|
async (payload) => {
|
package/src/AdkClient.ts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { SSEEventDecoder } from "assistant-stream/utils";
|
|
2
|
+
import { raceWithAbortSignal } from "@assistant-ui/core/internal";
|
|
2
3
|
import { contentToParts } from "./contentToParts";
|
|
3
4
|
import { parseAdkEventValue } from "./parseAdkEvent";
|
|
4
|
-
import { raceWithAbortSignal } from "./raceWithAbortSignal";
|
|
5
5
|
import { toAdkFunctionResponse } from "./toAdkFunctionResponse";
|
|
6
6
|
import { trimTrailingSlashes } from "./trimTrailingSlashes";
|
|
7
7
|
import type {
|
|
@@ -41,6 +41,12 @@ export type CreateAdkStreamOptions = {
|
|
|
41
41
|
| Record<string, string>
|
|
42
42
|
| (() => Record<string, string> | Promise<Record<string, string>>)
|
|
43
43
|
| undefined;
|
|
44
|
+
|
|
45
|
+
/** Maximum UTF-16 code units accepted in one SSE line. Defaults to 16 MiB. */
|
|
46
|
+
maxStreamLineLength?: number | undefined;
|
|
47
|
+
|
|
48
|
+
/** Maximum UTF-16 code units retained across one SSE event. Defaults to 16 MiB. */
|
|
49
|
+
maxStreamEventLength?: number | undefined;
|
|
44
50
|
};
|
|
45
51
|
|
|
46
52
|
/**
|
|
@@ -97,7 +103,8 @@ export function createAdkStream(
|
|
|
97
103
|
} else {
|
|
98
104
|
// Proxy mode: POST in parseAdkRequest-compatible format
|
|
99
105
|
url = options.api;
|
|
100
|
-
|
|
106
|
+
const { remoteId, externalId } = await config.initialize();
|
|
107
|
+
body = messagesToProxyBody(messages, config, externalId ?? remoteId);
|
|
101
108
|
}
|
|
102
109
|
|
|
103
110
|
const response = await fetch(url, {
|
|
@@ -114,7 +121,7 @@ export function createAdkStream(
|
|
|
114
121
|
}
|
|
115
122
|
|
|
116
123
|
validateEventStreamContentType(response);
|
|
117
|
-
yield* parseSSEResponse(response);
|
|
124
|
+
yield* parseSSEResponse(response, options);
|
|
118
125
|
};
|
|
119
126
|
}
|
|
120
127
|
|
|
@@ -209,8 +216,9 @@ function messagesToProxyBody(
|
|
|
209
216
|
checkpointId?: string | undefined;
|
|
210
217
|
stateDelta?: Record<string, unknown> | undefined;
|
|
211
218
|
},
|
|
219
|
+
sessionId: string,
|
|
212
220
|
): Record<string, unknown> {
|
|
213
|
-
const body: Record<string, unknown> = {};
|
|
221
|
+
const body: Record<string, unknown> = { sessionId };
|
|
214
222
|
|
|
215
223
|
if (config.runConfig != null) body.runConfig = config.runConfig;
|
|
216
224
|
if (config.checkpointId != null) body.checkpointId = config.checkpointId;
|
|
@@ -255,16 +263,26 @@ function messagesToProxyBody(
|
|
|
255
263
|
return body;
|
|
256
264
|
}
|
|
257
265
|
|
|
258
|
-
async function* parseSSEResponse(
|
|
266
|
+
async function* parseSSEResponse(
|
|
267
|
+
response: Response,
|
|
268
|
+
options: Pick<
|
|
269
|
+
CreateAdkStreamOptions,
|
|
270
|
+
"maxStreamLineLength" | "maxStreamEventLength"
|
|
271
|
+
>,
|
|
272
|
+
): AsyncGenerator<AdkEvent> {
|
|
259
273
|
if (!response.body) {
|
|
260
274
|
throw new Error("Expected ADK stream response body, received no body");
|
|
261
275
|
}
|
|
262
276
|
const reader = response.body.getReader();
|
|
263
277
|
const decoder = new TextDecoder();
|
|
264
|
-
const sseDecoder = new SSEEventDecoder({ trailing: "dispatch" });
|
|
265
278
|
|
|
266
279
|
let shouldCancel = true;
|
|
267
280
|
try {
|
|
281
|
+
const sseDecoder = new SSEEventDecoder({
|
|
282
|
+
trailing: "dispatch",
|
|
283
|
+
maxLineLength: options.maxStreamLineLength,
|
|
284
|
+
maxEventLength: options.maxStreamEventLength,
|
|
285
|
+
});
|
|
268
286
|
while (true) {
|
|
269
287
|
let result: ReadableStreamReadResult<Uint8Array>;
|
|
270
288
|
try {
|
package/src/AdkSessionAdapter.ts
CHANGED
|
@@ -5,12 +5,12 @@ import type {
|
|
|
5
5
|
RemoteThreadListResponse,
|
|
6
6
|
RemoteThreadMetadata,
|
|
7
7
|
} from "@assistant-ui/core";
|
|
8
|
+
import { raceWithAbortSignal } from "@assistant-ui/core/internal";
|
|
8
9
|
import { AdkEventAccumulator } from "./AdkEventAccumulator";
|
|
9
10
|
import { normalizeAdkPart } from "./normalizeAdkPart";
|
|
10
11
|
import { parseAdkEventValue } from "./parseAdkEvent";
|
|
11
12
|
import type { AdkMessage, AdkThreadSnapshot } from "./types";
|
|
12
13
|
import { trimTrailingSlashes } from "./trimTrailingSlashes";
|
|
13
|
-
import { raceWithAbortSignal } from "./raceWithAbortSignal";
|
|
14
14
|
|
|
15
15
|
export type AdkSessionAdapterOptions = {
|
|
16
16
|
/**
|
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
import { describe, expect, it, vi } from "vitest";
|
|
2
|
+
import { AdkThreadController } from "./AdkThreadController";
|
|
3
|
+
|
|
4
|
+
describe("AdkThreadController", () => {
|
|
5
|
+
it("selects staged messages through a parent in transcript order", () => {
|
|
6
|
+
const controller = new AdkThreadController();
|
|
7
|
+
const first = { id: "first", type: "human" as const, content: "one" };
|
|
8
|
+
const second = { id: "second", type: "human" as const, content: "two" };
|
|
9
|
+
const later = { id: "later", type: "human" as const, content: "three" };
|
|
10
|
+
controller.dispatch({
|
|
11
|
+
type: "staged.stage",
|
|
12
|
+
entry: { message: first, runConfig: { custom: { source: "first" } } },
|
|
13
|
+
});
|
|
14
|
+
controller.dispatch({
|
|
15
|
+
type: "staged.stage",
|
|
16
|
+
entry: { message: second, runConfig: { custom: { source: "second" } } },
|
|
17
|
+
});
|
|
18
|
+
controller.dispatch({
|
|
19
|
+
type: "staged.stage",
|
|
20
|
+
entry: { message: later, runConfig: undefined },
|
|
21
|
+
});
|
|
22
|
+
|
|
23
|
+
expect(controller.getStagedMessageCount()).toBe(3);
|
|
24
|
+
expect(
|
|
25
|
+
controller.getStagedRun("second", [
|
|
26
|
+
second,
|
|
27
|
+
{ id: "canonical", type: "ai", content: "reply" },
|
|
28
|
+
first,
|
|
29
|
+
later,
|
|
30
|
+
]),
|
|
31
|
+
).toEqual({
|
|
32
|
+
messages: [second],
|
|
33
|
+
runConfig: { custom: { source: "second" } },
|
|
34
|
+
});
|
|
35
|
+
expect(
|
|
36
|
+
controller.getStagedRun("later", [
|
|
37
|
+
second,
|
|
38
|
+
{ id: "canonical", type: "ai", content: "reply" },
|
|
39
|
+
first,
|
|
40
|
+
later,
|
|
41
|
+
])?.messages,
|
|
42
|
+
).toEqual([second, first, later]);
|
|
43
|
+
expect(controller.getStagedRun("canonical", [first, second])).toBeNull();
|
|
44
|
+
|
|
45
|
+
controller.dispatch({ type: "staged.unstage", ids: ["first", "second"] });
|
|
46
|
+
expect(controller.getStagedMessageCount()).toBe(1);
|
|
47
|
+
});
|
|
48
|
+
|
|
49
|
+
it("notifies for staging but not for unstaging an absent id", () => {
|
|
50
|
+
const controller = new AdkThreadController();
|
|
51
|
+
const listener = vi.fn();
|
|
52
|
+
controller.subscribe(listener);
|
|
53
|
+
|
|
54
|
+
controller.dispatch({
|
|
55
|
+
type: "staged.stage",
|
|
56
|
+
entry: {
|
|
57
|
+
message: { id: "first", type: "human", content: "one" },
|
|
58
|
+
runConfig: undefined,
|
|
59
|
+
},
|
|
60
|
+
});
|
|
61
|
+
expect(listener).toHaveBeenCalledTimes(1);
|
|
62
|
+
|
|
63
|
+
controller.dispatch({ type: "staged.unstage", ids: ["missing"] });
|
|
64
|
+
expect(listener).toHaveBeenCalledTimes(1);
|
|
65
|
+
});
|
|
66
|
+
|
|
67
|
+
it("publishes the reduced state synchronously to subscribers", () => {
|
|
68
|
+
const controller = new AdkThreadController();
|
|
69
|
+
const first = vi.fn(() => controller.getState());
|
|
70
|
+
const second = vi.fn();
|
|
71
|
+
const unsubscribeFirst = controller.subscribe(first);
|
|
72
|
+
controller.subscribe(second);
|
|
73
|
+
const initial = controller.getState();
|
|
74
|
+
const messages = [
|
|
75
|
+
{ id: "message-1", type: "human" as const, content: "hi" },
|
|
76
|
+
];
|
|
77
|
+
|
|
78
|
+
controller.dispatch({ type: "messages.set", messages });
|
|
79
|
+
|
|
80
|
+
expect(controller.getState()).not.toBe(initial);
|
|
81
|
+
expect(controller.getState().messages).toBe(messages);
|
|
82
|
+
expect(first).toHaveReturnedWith(controller.getState());
|
|
83
|
+
expect(second).toHaveBeenCalledTimes(1);
|
|
84
|
+
|
|
85
|
+
unsubscribeFirst();
|
|
86
|
+
controller.dispatch({ type: "messages.replaced", messages: [] });
|
|
87
|
+
expect(first).toHaveBeenCalledTimes(1);
|
|
88
|
+
expect(second).toHaveBeenCalledTimes(2);
|
|
89
|
+
});
|
|
90
|
+
});
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
import {
|
|
2
|
+
createAdkThreadState,
|
|
3
|
+
reduceAdkThreadState,
|
|
4
|
+
type AdkThreadAction,
|
|
5
|
+
} from "./adkThreadState";
|
|
6
|
+
import type { AdkMessage } from "./types";
|
|
7
|
+
|
|
8
|
+
export class AdkThreadController {
|
|
9
|
+
private state = createAdkThreadState();
|
|
10
|
+
private readonly listeners = new Set<() => void>();
|
|
11
|
+
|
|
12
|
+
public getState = () => this.state;
|
|
13
|
+
|
|
14
|
+
public getStagedMessageCount = () => this.state.stagedEntries.size;
|
|
15
|
+
|
|
16
|
+
public getStagedRun = (
|
|
17
|
+
parentId: string | null,
|
|
18
|
+
messages: AdkMessage[] = this.state.messages,
|
|
19
|
+
) => {
|
|
20
|
+
const entries = this.state.stagedEntries;
|
|
21
|
+
if (!parentId || !entries.has(parentId)) return null;
|
|
22
|
+
|
|
23
|
+
const staged: AdkMessage[] = [];
|
|
24
|
+
for (const message of messages) {
|
|
25
|
+
if (message.id && entries.has(message.id)) {
|
|
26
|
+
staged.push(entries.get(message.id)!.message);
|
|
27
|
+
}
|
|
28
|
+
if (message.id === parentId) break;
|
|
29
|
+
}
|
|
30
|
+
|
|
31
|
+
return { messages: staged, runConfig: entries.get(parentId)!.runConfig };
|
|
32
|
+
};
|
|
33
|
+
|
|
34
|
+
public subscribe = (listener: () => void) => {
|
|
35
|
+
this.listeners.add(listener);
|
|
36
|
+
return () => this.listeners.delete(listener);
|
|
37
|
+
};
|
|
38
|
+
|
|
39
|
+
public dispatch = (action: AdkThreadAction) => {
|
|
40
|
+
const nextState = reduceAdkThreadState(this.state, action);
|
|
41
|
+
if (nextState === this.state) return;
|
|
42
|
+
this.state = nextState;
|
|
43
|
+
for (const listener of this.listeners) listener();
|
|
44
|
+
};
|
|
45
|
+
}
|
|
@@ -0,0 +1,207 @@
|
|
|
1
|
+
import { describe, expect, it } from "vitest";
|
|
2
|
+
import { createAdkThreadState, reduceAdkThreadState } from "./adkThreadState";
|
|
3
|
+
import type { AdkMessage } from "./types";
|
|
4
|
+
|
|
5
|
+
const message: AdkMessage = {
|
|
6
|
+
id: "message-1",
|
|
7
|
+
type: "human",
|
|
8
|
+
content: "hello",
|
|
9
|
+
};
|
|
10
|
+
|
|
11
|
+
describe("reduceAdkThreadState", () => {
|
|
12
|
+
it("stages messages by id without changing the prior state", () => {
|
|
13
|
+
const initial = createAdkThreadState();
|
|
14
|
+
const first = {
|
|
15
|
+
message: { ...message },
|
|
16
|
+
runConfig: { custom: { model: "first" } },
|
|
17
|
+
};
|
|
18
|
+
const second = {
|
|
19
|
+
message: { ...message, id: "message-2" },
|
|
20
|
+
runConfig: { custom: { model: "second" } },
|
|
21
|
+
};
|
|
22
|
+
|
|
23
|
+
const staged = reduceAdkThreadState(initial, {
|
|
24
|
+
type: "staged.stage",
|
|
25
|
+
entry: first,
|
|
26
|
+
});
|
|
27
|
+
const next = reduceAdkThreadState(staged, {
|
|
28
|
+
type: "staged.stage",
|
|
29
|
+
entry: second,
|
|
30
|
+
});
|
|
31
|
+
|
|
32
|
+
expect(initial.stagedEntries.size).toBe(0);
|
|
33
|
+
expect(staged.stagedEntries?.size).toBe(1);
|
|
34
|
+
expect([...next.stagedEntries!]).toEqual([
|
|
35
|
+
["message-1", first],
|
|
36
|
+
["message-2", second],
|
|
37
|
+
]);
|
|
38
|
+
expect(staged.stagedEntries?.size).toBe(1);
|
|
39
|
+
});
|
|
40
|
+
|
|
41
|
+
it("unstages a promoted run and preserves remaining entries through snapshots", () => {
|
|
42
|
+
const entries = ["message-1", "message-2", "message-3"].map((id) => ({
|
|
43
|
+
message: { ...message, id },
|
|
44
|
+
runConfig: undefined,
|
|
45
|
+
}));
|
|
46
|
+
const staged = entries.reduce(
|
|
47
|
+
(state, entry) =>
|
|
48
|
+
reduceAdkThreadState(state, { type: "staged.stage", entry }),
|
|
49
|
+
createAdkThreadState(),
|
|
50
|
+
);
|
|
51
|
+
const promoted = reduceAdkThreadState(staged, {
|
|
52
|
+
type: "staged.unstage",
|
|
53
|
+
ids: ["message-1", "message-2"],
|
|
54
|
+
});
|
|
55
|
+
const snapshotted = reduceAdkThreadState(promoted, {
|
|
56
|
+
type: "snapshot.applied",
|
|
57
|
+
snapshot: { messages: [message] },
|
|
58
|
+
});
|
|
59
|
+
|
|
60
|
+
expect([...promoted.stagedEntries!.keys()]).toEqual(["message-3"]);
|
|
61
|
+
expect([...snapshotted.stagedEntries!.keys()]).toEqual(["message-3"]);
|
|
62
|
+
expect(staged.stagedEntries?.size).toBe(3);
|
|
63
|
+
expect(
|
|
64
|
+
reduceAdkThreadState(promoted, {
|
|
65
|
+
type: "staged.unstage",
|
|
66
|
+
ids: ["missing"],
|
|
67
|
+
}),
|
|
68
|
+
).toBe(promoted);
|
|
69
|
+
});
|
|
70
|
+
|
|
71
|
+
it("merges event deltas and metadata with previous state", () => {
|
|
72
|
+
const previous = {
|
|
73
|
+
...createAdkThreadState(),
|
|
74
|
+
stagedEntries: new Map([
|
|
75
|
+
["message-1", { message, runConfig: undefined }],
|
|
76
|
+
]),
|
|
77
|
+
stateDelta: { retained: 1, replaced: "old" },
|
|
78
|
+
artifactDelta: { retained: 2, replaced: 3 },
|
|
79
|
+
messageMetadata: new Map([
|
|
80
|
+
["retained", { groundingMetadata: "old" }],
|
|
81
|
+
["replaced", { citationMetadata: "old" }],
|
|
82
|
+
]),
|
|
83
|
+
};
|
|
84
|
+
const published = {
|
|
85
|
+
...createAdkThreadState(),
|
|
86
|
+
messages: [message],
|
|
87
|
+
stateDelta: { replaced: "new" },
|
|
88
|
+
artifactDelta: { replaced: 4 },
|
|
89
|
+
agentInfo: { name: "agent" },
|
|
90
|
+
longRunningToolIds: ["tool-1"],
|
|
91
|
+
escalated: true,
|
|
92
|
+
messageMetadata: new Map([["replaced", { citationMetadata: "new" }]]),
|
|
93
|
+
};
|
|
94
|
+
const { stagedEntries: _stagedEntries, ...publishedState } = published;
|
|
95
|
+
|
|
96
|
+
const next = reduceAdkThreadState(previous, {
|
|
97
|
+
type: "event.published",
|
|
98
|
+
state: publishedState,
|
|
99
|
+
});
|
|
100
|
+
|
|
101
|
+
expect(next).toMatchObject({
|
|
102
|
+
messages: [message],
|
|
103
|
+
stateDelta: { retained: 1, replaced: "new" },
|
|
104
|
+
artifactDelta: { retained: 2, replaced: 4 },
|
|
105
|
+
agentInfo: { name: "agent" },
|
|
106
|
+
longRunningToolIds: ["tool-1"],
|
|
107
|
+
escalated: true,
|
|
108
|
+
});
|
|
109
|
+
expect([...next.messageMetadata]).toEqual([
|
|
110
|
+
["retained", { groundingMetadata: "old" }],
|
|
111
|
+
["replaced", { citationMetadata: "new" }],
|
|
112
|
+
]);
|
|
113
|
+
expect(next.stagedEntries).toBe(previous.stagedEntries);
|
|
114
|
+
expect(previous.stateDelta).toEqual({ retained: 1, replaced: "old" });
|
|
115
|
+
expect(previous.messageMetadata.get("replaced")).toEqual({
|
|
116
|
+
citationMetadata: "old",
|
|
117
|
+
});
|
|
118
|
+
|
|
119
|
+
const withoutMetadata = reduceAdkThreadState(next, {
|
|
120
|
+
type: "event.published",
|
|
121
|
+
state: { ...publishedState, messageMetadata: new Map() },
|
|
122
|
+
});
|
|
123
|
+
expect(withoutMetadata.messageMetadata).toBe(next.messageMetadata);
|
|
124
|
+
});
|
|
125
|
+
|
|
126
|
+
it("swaps a snapshot and clears omitted fields", () => {
|
|
127
|
+
const previous = {
|
|
128
|
+
...createAdkThreadState(),
|
|
129
|
+
stateDelta: { stale: true },
|
|
130
|
+
artifactDelta: { stale: 1 },
|
|
131
|
+
agentInfo: { name: "stale-agent", branch: "stale" },
|
|
132
|
+
longRunningToolIds: ["stale"],
|
|
133
|
+
toolConfirmations: [
|
|
134
|
+
{
|
|
135
|
+
toolCallId: "stale-tool",
|
|
136
|
+
toolName: "approve",
|
|
137
|
+
args: {},
|
|
138
|
+
hint: "stale",
|
|
139
|
+
confirmed: false,
|
|
140
|
+
},
|
|
141
|
+
],
|
|
142
|
+
authRequests: [{ toolCallId: "stale-auth", authConfig: {} }],
|
|
143
|
+
escalated: true,
|
|
144
|
+
messageMetadata: new Map([["stale", { groundingMetadata: "stale" }]]),
|
|
145
|
+
};
|
|
146
|
+
const snapshot = {
|
|
147
|
+
messages: [message],
|
|
148
|
+
stateDelta: { loaded: true },
|
|
149
|
+
messageMetadata: new Map([["message-1", { usageMetadata: 1 }]]),
|
|
150
|
+
};
|
|
151
|
+
|
|
152
|
+
const next = reduceAdkThreadState(previous, {
|
|
153
|
+
type: "snapshot.applied",
|
|
154
|
+
snapshot,
|
|
155
|
+
});
|
|
156
|
+
|
|
157
|
+
expect(next).toEqual({
|
|
158
|
+
messages: snapshot.messages,
|
|
159
|
+
stateDelta: snapshot.stateDelta,
|
|
160
|
+
agentInfo: {},
|
|
161
|
+
longRunningToolIds: [],
|
|
162
|
+
artifactDelta: {},
|
|
163
|
+
toolConfirmations: [],
|
|
164
|
+
authRequests: [],
|
|
165
|
+
escalated: false,
|
|
166
|
+
messageMetadata: snapshot.messageMetadata,
|
|
167
|
+
stagedEntries: previous.stagedEntries,
|
|
168
|
+
});
|
|
169
|
+
});
|
|
170
|
+
|
|
171
|
+
it("replaces messages and clears per-turn state while retaining thread deltas", () => {
|
|
172
|
+
const previous = {
|
|
173
|
+
...createAdkThreadState(),
|
|
174
|
+
stateDelta: { retained: true },
|
|
175
|
+
artifactDelta: { file: 1 },
|
|
176
|
+
agentInfo: { name: "agent" },
|
|
177
|
+
longRunningToolIds: ["tool-1"],
|
|
178
|
+
toolConfirmations: [
|
|
179
|
+
{
|
|
180
|
+
toolCallId: "tool-1",
|
|
181
|
+
toolName: "search",
|
|
182
|
+
args: {},
|
|
183
|
+
hint: "approve",
|
|
184
|
+
confirmed: false,
|
|
185
|
+
},
|
|
186
|
+
],
|
|
187
|
+
authRequests: [{ toolCallId: "tool-2", authConfig: {} }],
|
|
188
|
+
escalated: true,
|
|
189
|
+
messageMetadata: new Map([["old", { usageMetadata: 1 }]]),
|
|
190
|
+
};
|
|
191
|
+
|
|
192
|
+
const next = reduceAdkThreadState(previous, {
|
|
193
|
+
type: "messages.replaced",
|
|
194
|
+
messages: [message],
|
|
195
|
+
});
|
|
196
|
+
|
|
197
|
+
expect(next.messages).toEqual([message]);
|
|
198
|
+
expect(next.longRunningToolIds).toEqual([]);
|
|
199
|
+
expect(next.toolConfirmations).toEqual([]);
|
|
200
|
+
expect(next.authRequests).toEqual([]);
|
|
201
|
+
expect(next.escalated).toBe(false);
|
|
202
|
+
expect(next.messageMetadata.size).toBe(0);
|
|
203
|
+
expect(next.stateDelta).toBe(previous.stateDelta);
|
|
204
|
+
expect(next.artifactDelta).toBe(previous.artifactDelta);
|
|
205
|
+
expect(next.agentInfo).toBe(previous.agentInfo);
|
|
206
|
+
});
|
|
207
|
+
});
|
|
@@ -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,
|