@assistant-ui/react-a2a 0.2.22 → 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.
@@ -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
+ });
@@ -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
- // Create or reuse client
28
- const clientRef = useRef<A2AClient | null>(null);
29
- if (!clientRef.current) {
30
- if (options.client) {
31
- clientRef.current = options.client;
32
- } else if (options.baseUrl) {
33
- clientRef.current = new A2AClient({
34
- baseUrl: options.baseUrl,
35
- basePath: options.basePath,
36
- tenant: options.tenant,
37
- headers: options.headers,
38
- extensions: options.extensions,
39
- fetchOptions: options.fetchOptions,
40
- });
41
- } else {
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
- const client = clientRef.current;
46
-
47
- // Create or reuse core
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
- messages: core.getMessages(),
141
+ messageRepository: core.getMessageRepository(),
123
142
  isRunning: core.isRunning(),
124
143
  extras: a2aExtras.provide({
125
144
  task: core.getTask(),