@assistant-ui/react-a2a 0.2.21 → 0.2.23
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/dist/A2AClient.d.ts +8 -5
- package/dist/A2AClient.d.ts.map +1 -1
- package/dist/A2AClient.js +85 -24
- package/dist/A2AClient.js.map +1 -1
- package/dist/A2AThreadRuntimeCore.d.ts +15 -4
- package/dist/A2AThreadRuntimeCore.d.ts.map +1 -1
- package/dist/A2AThreadRuntimeCore.js +162 -45
- package/dist/A2AThreadRuntimeCore.js.map +1 -1
- package/dist/a2aExtras.d.ts +0 -1
- package/dist/a2aExtras.d.ts.map +1 -1
- package/dist/conversions.d.ts +6 -4
- package/dist/conversions.d.ts.map +1 -1
- package/dist/conversions.js +32 -3
- package/dist/conversions.js.map +1 -1
- package/dist/hooks.d.ts +0 -1
- package/dist/hooks.d.ts.map +1 -1
- package/dist/types.d.ts +27 -14
- package/dist/types.d.ts.map +1 -1
- package/dist/useA2ARuntime.d.ts +0 -1
- package/dist/useA2ARuntime.d.ts.map +1 -1
- package/dist/useA2ARuntime.js +30 -18
- package/dist/useA2ARuntime.js.map +1 -1
- package/package.json +11 -6
- package/src/A2AClient.test.ts +321 -4
- package/src/A2AClient.ts +158 -49
- package/src/A2AThreadRuntimeCore.test.ts +458 -0
- package/src/A2AThreadRuntimeCore.ts +215 -57
- package/src/conversions.test.ts +109 -2
- package/src/conversions.ts +46 -4
- package/src/useA2ARuntime.test.tsx +142 -0
- package/src/useA2ARuntime.ts +55 -36
|
@@ -0,0 +1,142 @@
|
|
|
1
|
+
// @vitest-environment jsdom
|
|
2
|
+
|
|
3
|
+
import { act, renderHook, waitFor } from "@testing-library/react";
|
|
4
|
+
import { afterEach, describe, expect, it, vi } from "vitest";
|
|
5
|
+
import type { A2AClient } from "./A2AClient";
|
|
6
|
+
import { useA2ARuntime } from "./useA2ARuntime";
|
|
7
|
+
|
|
8
|
+
const createMockClient = (waitForAbort = false) => {
|
|
9
|
+
let streamSignal: AbortSignal | undefined;
|
|
10
|
+
const getAgentCard = vi.fn().mockResolvedValue(undefined);
|
|
11
|
+
const streamMessage = vi.fn(
|
|
12
|
+
(
|
|
13
|
+
_message: unknown,
|
|
14
|
+
_configuration: unknown,
|
|
15
|
+
_metadata: unknown,
|
|
16
|
+
signal?: AbortSignal,
|
|
17
|
+
): AsyncIterable<never> => {
|
|
18
|
+
streamSignal = signal;
|
|
19
|
+
return {
|
|
20
|
+
async *[Symbol.asyncIterator]() {
|
|
21
|
+
if (!waitForAbort || !signal) return;
|
|
22
|
+
await new Promise<void>((resolve) => {
|
|
23
|
+
if (signal.aborted) {
|
|
24
|
+
resolve();
|
|
25
|
+
} else {
|
|
26
|
+
signal.addEventListener("abort", () => resolve(), {
|
|
27
|
+
once: true,
|
|
28
|
+
});
|
|
29
|
+
}
|
|
30
|
+
});
|
|
31
|
+
},
|
|
32
|
+
};
|
|
33
|
+
},
|
|
34
|
+
);
|
|
35
|
+
|
|
36
|
+
return {
|
|
37
|
+
client: {
|
|
38
|
+
getAgentCard,
|
|
39
|
+
streamMessage,
|
|
40
|
+
} as unknown as A2AClient,
|
|
41
|
+
getAgentCard,
|
|
42
|
+
streamMessage,
|
|
43
|
+
get streamSignal() {
|
|
44
|
+
return streamSignal;
|
|
45
|
+
},
|
|
46
|
+
};
|
|
47
|
+
};
|
|
48
|
+
|
|
49
|
+
afterEach(() => {
|
|
50
|
+
vi.unstubAllGlobals();
|
|
51
|
+
});
|
|
52
|
+
|
|
53
|
+
describe("useA2ARuntime", () => {
|
|
54
|
+
it("switches provided clients and aborts the previous client run", async () => {
|
|
55
|
+
const first = createMockClient(true);
|
|
56
|
+
const second = createMockClient();
|
|
57
|
+
const { result, rerender } = renderHook(
|
|
58
|
+
({ client }) => useA2ARuntime({ client }),
|
|
59
|
+
{ initialProps: { client: first.client } },
|
|
60
|
+
);
|
|
61
|
+
|
|
62
|
+
await waitFor(() => expect(first.getAgentCard).toHaveBeenCalledOnce());
|
|
63
|
+
|
|
64
|
+
act(() => {
|
|
65
|
+
result.current.thread.append({
|
|
66
|
+
role: "user",
|
|
67
|
+
content: [{ type: "text", text: "first" }],
|
|
68
|
+
});
|
|
69
|
+
});
|
|
70
|
+
await waitFor(() => expect(first.streamMessage).toHaveBeenCalledOnce());
|
|
71
|
+
|
|
72
|
+
rerender({ client: second.client });
|
|
73
|
+
|
|
74
|
+
await waitFor(() => expect(first.streamSignal?.aborted).toBe(true));
|
|
75
|
+
await waitFor(() => expect(second.getAgentCard).toHaveBeenCalledOnce());
|
|
76
|
+
|
|
77
|
+
act(() => {
|
|
78
|
+
result.current.thread.append({
|
|
79
|
+
role: "user",
|
|
80
|
+
content: [{ type: "text", text: "second" }],
|
|
81
|
+
});
|
|
82
|
+
});
|
|
83
|
+
|
|
84
|
+
await waitFor(() => expect(second.streamMessage).toHaveBeenCalledOnce());
|
|
85
|
+
expect(first.streamMessage).toHaveBeenCalledOnce();
|
|
86
|
+
});
|
|
87
|
+
|
|
88
|
+
it("uses current headers without recreating the managed client", async () => {
|
|
89
|
+
const fetchMock = vi.fn(
|
|
90
|
+
async (input: string | URL | Request, _init?: RequestInit) => {
|
|
91
|
+
const url =
|
|
92
|
+
typeof input === "string"
|
|
93
|
+
? input
|
|
94
|
+
: input instanceof URL
|
|
95
|
+
? input.href
|
|
96
|
+
: input.url;
|
|
97
|
+
if (url.endsWith("/.well-known/agent-card.json")) {
|
|
98
|
+
return new Response(
|
|
99
|
+
JSON.stringify({ capabilities: { streaming: true } }),
|
|
100
|
+
{
|
|
101
|
+
status: 200,
|
|
102
|
+
headers: { "Content-Type": "application/json" },
|
|
103
|
+
},
|
|
104
|
+
);
|
|
105
|
+
}
|
|
106
|
+
return new Response("", {
|
|
107
|
+
status: 200,
|
|
108
|
+
headers: { "Content-Type": "text/event-stream" },
|
|
109
|
+
});
|
|
110
|
+
},
|
|
111
|
+
);
|
|
112
|
+
vi.stubGlobal("fetch", fetchMock);
|
|
113
|
+
|
|
114
|
+
const { result, rerender } = renderHook(
|
|
115
|
+
({ token }) =>
|
|
116
|
+
useA2ARuntime({
|
|
117
|
+
baseUrl: "https://agent.test",
|
|
118
|
+
headers: { Authorization: `Bearer ${token}` },
|
|
119
|
+
extensions: ["urn:example"],
|
|
120
|
+
fetchOptions: { credentials: "include" },
|
|
121
|
+
}),
|
|
122
|
+
{ initialProps: { token: "first" } },
|
|
123
|
+
);
|
|
124
|
+
|
|
125
|
+
await waitFor(() => expect(fetchMock).toHaveBeenCalledOnce());
|
|
126
|
+
|
|
127
|
+
rerender({ token: "second" });
|
|
128
|
+
|
|
129
|
+
act(() => {
|
|
130
|
+
result.current.thread.append({
|
|
131
|
+
role: "user",
|
|
132
|
+
content: [{ type: "text", text: "hello" }],
|
|
133
|
+
});
|
|
134
|
+
});
|
|
135
|
+
|
|
136
|
+
await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(2));
|
|
137
|
+
const streamRequest = fetchMock.mock.calls[1]!;
|
|
138
|
+
expect(streamRequest[1]?.headers).toMatchObject({
|
|
139
|
+
Authorization: "Bearer second",
|
|
140
|
+
});
|
|
141
|
+
});
|
|
142
|
+
});
|
package/src/useA2ARuntime.ts
CHANGED
|
@@ -12,11 +12,28 @@ import type {
|
|
|
12
12
|
ExternalStoreAdapter,
|
|
13
13
|
ThreadMessage,
|
|
14
14
|
} from "@assistant-ui/core";
|
|
15
|
-
import { A2AClient } from "./A2AClient";
|
|
15
|
+
import { A2AClient, type A2AClientOptions } from "./A2AClient";
|
|
16
16
|
import { A2AThreadRuntimeCore } from "./A2AThreadRuntimeCore";
|
|
17
17
|
import { a2aExtras } from "./a2aExtras";
|
|
18
18
|
import type { UseA2ARuntimeOptions } from "./types";
|
|
19
19
|
|
|
20
|
+
type ManagedA2AClientOptions = Omit<A2AClientOptions, "headers">;
|
|
21
|
+
|
|
22
|
+
const serializeManagedClientOptions = (
|
|
23
|
+
options: ManagedA2AClientOptions,
|
|
24
|
+
): string => {
|
|
25
|
+
const fetchOptions = Object.fromEntries(
|
|
26
|
+
Object.entries(options.fetchOptions ?? {}).sort(([a], [b]) =>
|
|
27
|
+
a.localeCompare(b),
|
|
28
|
+
),
|
|
29
|
+
);
|
|
30
|
+
|
|
31
|
+
return JSON.stringify({
|
|
32
|
+
...options,
|
|
33
|
+
fetchOptions,
|
|
34
|
+
});
|
|
35
|
+
};
|
|
36
|
+
|
|
20
37
|
export function useA2ARuntime(options: UseA2ARuntimeOptions): AssistantRuntime {
|
|
21
38
|
const [_version, setVersion] = useState(0);
|
|
22
39
|
const notifyUpdate = useCallback(() => setVersion((v) => v + 1), []);
|
|
@@ -24,44 +41,46 @@ export function useA2ARuntime(options: UseA2ARuntimeOptions): AssistantRuntime {
|
|
|
24
41
|
const historyAdapter = options.adapters?.history ?? runtimeAdapters?.history;
|
|
25
42
|
const threadListAdapter = options.adapters?.threadList;
|
|
26
43
|
|
|
27
|
-
|
|
28
|
-
|
|
29
|
-
|
|
30
|
-
|
|
31
|
-
|
|
32
|
-
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
|
|
38
|
-
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
|
|
44
|
+
const headersRef = useRef(options.headers);
|
|
45
|
+
headersRef.current = options.headers;
|
|
46
|
+
const resolveHeaders = useCallback(() => {
|
|
47
|
+
const headers = headersRef.current;
|
|
48
|
+
return typeof headers === "function" ? headers() : (headers ?? {});
|
|
49
|
+
}, []);
|
|
50
|
+
|
|
51
|
+
const managedClientOptionsKey = options.client
|
|
52
|
+
? null
|
|
53
|
+
: options.baseUrl
|
|
54
|
+
? serializeManagedClientOptions({
|
|
55
|
+
baseUrl: options.baseUrl,
|
|
56
|
+
basePath: options.basePath,
|
|
57
|
+
tenant: options.tenant,
|
|
58
|
+
extensions: options.extensions,
|
|
59
|
+
fetchOptions: options.fetchOptions,
|
|
60
|
+
})
|
|
61
|
+
: null;
|
|
62
|
+
|
|
63
|
+
const client = useMemo(() => {
|
|
64
|
+
if (options.client) return options.client;
|
|
65
|
+
if (!managedClientOptionsKey) {
|
|
42
66
|
throw new Error("useA2ARuntime requires either `client` or `baseUrl`");
|
|
43
67
|
}
|
|
44
|
-
|
|
45
|
-
|
|
46
|
-
|
|
47
|
-
|
|
48
|
-
const coreRef = useRef<A2AThreadRuntimeCore | null>(null);
|
|
49
|
-
if (!coreRef.current) {
|
|
50
|
-
coreRef.current = new A2AThreadRuntimeCore({
|
|
51
|
-
client,
|
|
52
|
-
contextId: options.contextId,
|
|
53
|
-
configuration: options.configuration,
|
|
54
|
-
...(options.onError && { onError: options.onError }),
|
|
55
|
-
...(options.onCancel && { onCancel: options.onCancel }),
|
|
56
|
-
...(options.onArtifactComplete && {
|
|
57
|
-
onArtifactComplete: options.onArtifactComplete,
|
|
58
|
-
}),
|
|
59
|
-
...(historyAdapter && { history: historyAdapter }),
|
|
60
|
-
notifyUpdate,
|
|
68
|
+
|
|
69
|
+
return new A2AClient({
|
|
70
|
+
...(JSON.parse(managedClientOptionsKey) as ManagedA2AClientOptions),
|
|
71
|
+
headers: resolveHeaders,
|
|
61
72
|
});
|
|
62
|
-
}
|
|
73
|
+
}, [managedClientOptionsKey, options.client, resolveHeaders]);
|
|
74
|
+
|
|
75
|
+
const core = useMemo(
|
|
76
|
+
() =>
|
|
77
|
+
new A2AThreadRuntimeCore({
|
|
78
|
+
client,
|
|
79
|
+
notifyUpdate,
|
|
80
|
+
}),
|
|
81
|
+
[client, notifyUpdate],
|
|
82
|
+
);
|
|
63
83
|
|
|
64
|
-
const core = coreRef.current;
|
|
65
84
|
core.updateOptions({
|
|
66
85
|
client,
|
|
67
86
|
contextId: options.contextId,
|
|
@@ -119,7 +138,7 @@ export function useA2ARuntime(options: UseA2ARuntimeOptions): AssistantRuntime {
|
|
|
119
138
|
return {
|
|
120
139
|
...shared,
|
|
121
140
|
isLoading: core.isLoading,
|
|
122
|
-
|
|
141
|
+
messageRepository: core.getMessageRepository(),
|
|
123
142
|
isRunning: core.isRunning(),
|
|
124
143
|
extras: a2aExtras.provide({
|
|
125
144
|
task: core.getTask(),
|