@assistant-ui/react-a2a 0.2.33 → 0.2.35

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.
@@ -1,7 +1,12 @@
1
1
  import { describe, it, expect, vi, beforeEach, afterEach } from "vitest";
2
2
  import { A2AThreadRuntimeCore } from "./A2AThreadRuntimeCore";
3
3
  import type { A2AClient } from "./A2AClient";
4
- import type { A2AMessage, A2AStreamEvent, A2ATask } from "./types";
4
+ import type {
5
+ A2AAgentCard,
6
+ A2AMessage,
7
+ A2AStreamEvent,
8
+ A2ATask,
9
+ } from "./types";
5
10
  import type { AppendMessage, ThreadMessage } from "@assistant-ui/core";
6
11
 
7
12
  // --- Mock client factory ---
@@ -791,6 +796,147 @@ describe("A2AThreadRuntimeCore", () => {
791
796
  // --- Sync (non-streaming) fallback ---
792
797
 
793
798
  describe("sync fallback", () => {
799
+ it("retries agent card discovery after a transient failure", async () => {
800
+ let now = 1_000;
801
+ vi.spyOn(Date, "now").mockImplementation(() => now);
802
+ let resolveRecovery!: (card: A2AAgentCard) => void;
803
+ const transportOrder: string[] = [];
804
+ const getAgentCard = vi
805
+ .fn()
806
+ .mockRejectedValueOnce(new Error("temporary failure"))
807
+ .mockImplementationOnce(
808
+ () =>
809
+ new Promise<A2AAgentCard>((resolve) => {
810
+ resolveRecovery = resolve;
811
+ }),
812
+ );
813
+ const sendMessage = vi.fn().mockImplementation(async () => {
814
+ transportOrder.push("sync");
815
+ return {
816
+ id: "t1",
817
+ status: { state: "completed" },
818
+ } satisfies A2ATask;
819
+ });
820
+ const streamMessage = vi.fn().mockImplementation(async function* () {
821
+ transportOrder.push("stream");
822
+ yield statusUpdateEvent("completed");
823
+ });
824
+ const core = createCore({ getAgentCard, sendMessage, streamMessage });
825
+
826
+ await core.append(createUserAppendMessage("First"));
827
+ now += 5_000;
828
+ await core.append(createUserAppendMessage("Second"));
829
+ resolveRecovery({
830
+ name: "Agent",
831
+ capabilities: { streaming: false },
832
+ } as A2AAgentCard);
833
+ await vi.waitFor(() => expect(core.getAgentCard()).toBeDefined());
834
+ await core.append(createUserAppendMessage("Third"));
835
+
836
+ expect(getAgentCard).toHaveBeenCalledTimes(2);
837
+ expect(streamMessage).toHaveBeenCalledTimes(2);
838
+ expect(sendMessage).toHaveBeenCalledOnce();
839
+ expect(transportOrder).toEqual(["stream", "stream", "sync"]);
840
+ });
841
+
842
+ it("does not repeat failed discovery during the retry delay", async () => {
843
+ vi.spyOn(Date, "now").mockReturnValue(1_000);
844
+ const getAgentCard = vi.fn().mockRejectedValue(new Error("unavailable"));
845
+ const streamMessage = vi.fn().mockImplementation(async function* () {
846
+ yield statusUpdateEvent("completed");
847
+ });
848
+ const core = createCore({ getAgentCard, streamMessage });
849
+
850
+ await core.append(createUserAppendMessage("First"));
851
+ await core.append(createUserAppendMessage("Second"));
852
+
853
+ expect(getAgentCard).toHaveBeenCalledOnce();
854
+ expect(streamMessage).toHaveBeenCalledTimes(2);
855
+ });
856
+
857
+ it("backs off persistent failures without blocking later sends", async () => {
858
+ let now = 1_000;
859
+ vi.spyOn(Date, "now").mockImplementation(() => now);
860
+ let rejectSecond!: (error: Error) => void;
861
+ const getAgentCard = vi
862
+ .fn()
863
+ .mockRejectedValueOnce(new Error("first failure"))
864
+ .mockImplementationOnce(
865
+ () =>
866
+ new Promise<A2AAgentCard>((_resolve, reject) => {
867
+ rejectSecond = reject;
868
+ }),
869
+ )
870
+ .mockRejectedValue(new Error("still unavailable"));
871
+ const streamMessage = vi.fn().mockImplementation(async function* () {
872
+ yield statusUpdateEvent("completed");
873
+ });
874
+ const core = createCore({ getAgentCard, streamMessage });
875
+
876
+ await core.append(createUserAppendMessage("First"));
877
+ now = 6_000;
878
+ await core.append(createUserAppendMessage("Second"));
879
+ expect(streamMessage).toHaveBeenCalledTimes(2);
880
+ rejectSecond(new Error("second failure"));
881
+ await vi.waitFor(() => expect(getAgentCard).toHaveBeenCalledTimes(2));
882
+
883
+ now = 15_999;
884
+ await core.append(createUserAppendMessage("Third"));
885
+ expect(getAgentCard).toHaveBeenCalledTimes(2);
886
+
887
+ now = 16_000;
888
+ await core.append(createUserAppendMessage("Fourth"));
889
+ expect(getAgentCard).toHaveBeenCalledTimes(3);
890
+ expect(streamMessage).toHaveBeenCalledTimes(4);
891
+ });
892
+
893
+ it("waits for agent capabilities before choosing the first send method", async () => {
894
+ let resolveAgentCard!: (value: A2AAgentCard) => void;
895
+ const getAgentCard = vi.fn(
896
+ () =>
897
+ new Promise<A2AAgentCard>((resolve) => {
898
+ resolveAgentCard = resolve;
899
+ }),
900
+ );
901
+ const sendMessage = vi.fn().mockResolvedValue({
902
+ id: "t1",
903
+ status: { state: "completed" },
904
+ } satisfies A2ATask);
905
+ const streamMessage = vi.fn();
906
+ const core = createCore({ getAgentCard, sendMessage, streamMessage });
907
+
908
+ const run = core.append(createUserAppendMessage("Hello"));
909
+ expect(sendMessage).not.toHaveBeenCalled();
910
+ expect(streamMessage).not.toHaveBeenCalled();
911
+
912
+ resolveAgentCard({
913
+ name: "Agent",
914
+ capabilities: { streaming: false },
915
+ } as A2AAgentCard);
916
+ await run;
917
+
918
+ expect(getAgentCard).toHaveBeenCalledOnce();
919
+ expect(sendMessage).toHaveBeenCalledOnce();
920
+ expect(streamMessage).not.toHaveBeenCalled();
921
+ });
922
+
923
+ it("stops waiting for agent capabilities when the run is cancelled", async () => {
924
+ const getAgentCard = vi.fn(() => new Promise<A2AAgentCard>(() => {}));
925
+ const sendMessage = vi.fn();
926
+ const streamMessage = vi.fn();
927
+ const core = createCore({ getAgentCard, sendMessage, streamMessage });
928
+
929
+ const run = core.append(createUserAppendMessage("Hello"));
930
+ expect(core.isRunning()).toBe(true);
931
+
932
+ await core.cancel();
933
+ await run;
934
+
935
+ expect(core.isRunning()).toBe(false);
936
+ expect(sendMessage).not.toHaveBeenCalled();
937
+ expect(streamMessage).not.toHaveBeenCalled();
938
+ });
939
+
794
940
  it("uses sendMessage when streaming is false in agent card", async () => {
795
941
  const sendMessage = vi.fn().mockResolvedValue({
796
942
  id: "t1",
@@ -25,6 +25,7 @@ import type {
25
25
  A2ATaskArtifactUpdateEvent,
26
26
  A2ATaskStatusUpdateEvent,
27
27
  } from "./types";
28
+
28
29
  import {
29
30
  a2aMessageToContent,
30
31
  isTerminalTaskState,
@@ -32,6 +33,9 @@ import {
32
33
  taskStateToMessageStatus,
33
34
  } from "./conversions";
34
35
 
36
+ const INITIAL_AGENT_CARD_RETRY_DELAY_MS = 5_000;
37
+ const MAX_AGENT_CARD_RETRY_DELAY_MS = 5 * 60_000;
38
+
35
39
  export type A2AThreadRuntimeCoreOptions = {
36
40
  client: A2AClient;
37
41
  contextId?: string | undefined;
@@ -93,6 +97,9 @@ export class A2AThreadRuntimeCore {
93
97
  private _loadPromise: Promise<void> | undefined;
94
98
  private _loadRequested = false;
95
99
  private _agentCardPromise: Promise<void> | undefined;
100
+ private _agentCardRetryAfter = 0;
101
+ private _agentCardRetryDelay = INITIAL_AGENT_CARD_RETRY_DELAY_MS;
102
+ private _agentCardDiscoveryFailed = false;
96
103
 
97
104
  private lastOptionsContextId: string | undefined;
98
105
 
@@ -197,24 +204,62 @@ export class A2AThreadRuntimeCore {
197
204
  return this._isLoading;
198
205
  }
199
206
 
200
- __internal_load(): Promise<void> {
201
- this._loadRequested = true;
202
- this._agentCardPromise ??= this.client
203
- .getAgentCard()
204
- .then((agentCard) => {
207
+ private loadAgentCard(): Promise<void> {
208
+ if (Date.now() < this._agentCardRetryAfter) return Promise.resolve();
209
+
210
+ this._agentCardPromise ??= this.client.getAgentCard().then(
211
+ (agentCard) => {
205
212
  this.agentCardValue = agentCard;
213
+ this._agentCardRetryAfter = 0;
214
+ this._agentCardRetryDelay = INITIAL_AGENT_CARD_RETRY_DELAY_MS;
215
+ this._agentCardDiscoveryFailed = false;
206
216
  this.notifyUpdate();
207
- })
208
- .catch(() => undefined);
217
+ },
218
+ () => {
219
+ this._agentCardDiscoveryFailed = true;
220
+ this._agentCardRetryAfter = Date.now() + this._agentCardRetryDelay;
221
+ this._agentCardRetryDelay = Math.min(
222
+ this._agentCardRetryDelay * 2,
223
+ MAX_AGENT_CARD_RETRY_DELAY_MS,
224
+ );
225
+ this._agentCardPromise = undefined;
226
+ },
227
+ );
228
+ return this._agentCardPromise;
229
+ }
230
+
231
+ private async waitForAgentCard(signal: AbortSignal): Promise<boolean> {
232
+ const shouldWait = !this._agentCardDiscoveryFailed;
233
+ const load = this.loadAgentCard();
234
+ if (signal.aborted) return false;
235
+ if (!shouldWait) return true;
236
+
237
+ let onAbort!: () => void;
238
+ const abort = new Promise<void>((resolve) => {
239
+ onAbort = resolve;
240
+ signal.addEventListener("abort", onAbort, { once: true });
241
+ });
242
+
243
+ try {
244
+ await Promise.race([load, abort]);
245
+ } finally {
246
+ signal.removeEventListener("abort", onAbort);
247
+ }
248
+ return !signal.aborted;
249
+ }
250
+
251
+ __internal_load(): Promise<void> {
252
+ this._loadRequested = true;
253
+ const agentCardPromise = this.loadAgentCard();
209
254
 
210
255
  if (this._loadPromise) return this._loadPromise;
211
- if (!this.history) return this._agentCardPromise;
256
+ if (!this.history) return agentCardPromise;
212
257
 
213
258
  this._isLoading = true;
214
259
 
215
260
  const historyPromise = this.history.load();
216
261
 
217
- this._loadPromise = Promise.all([historyPromise, this._agentCardPromise])
262
+ this._loadPromise = Promise.all([historyPromise, agentCardPromise])
218
263
  .then(([repo]) => {
219
264
  if (repo) {
220
265
  this.session.applyExternalMessageRepository(repo);
@@ -407,11 +452,11 @@ export class A2AThreadRuntimeCore {
407
452
 
408
453
  this.setRunning(true);
409
454
 
410
- // Check if agent supports streaming; fall back to sync sendMessage if not
411
- const supportsStreaming =
412
- this.agentCardValue?.capabilities?.streaming !== false;
413
-
414
455
  try {
456
+ if (!(await this.waitForAgentCard(abortController.signal))) return;
457
+
458
+ const supportsStreaming =
459
+ this.agentCardValue?.capabilities?.streaming !== false;
415
460
  if (supportsStreaming) {
416
461
  await this.runStreaming(a2aMessage, assistantId, abortController);
417
462
  } else {
@@ -1,8 +1,14 @@
1
1
  // @vitest-environment jsdom
2
2
 
3
3
  import { act, renderHook, waitFor } from "@testing-library/react";
4
- import { startTransition, Suspense, type PropsWithChildren } from "react";
4
+ import {
5
+ startTransition,
6
+ Suspense,
7
+ useState,
8
+ type PropsWithChildren,
9
+ } from "react";
5
10
  import { afterEach, describe, expect, it, vi } from "vitest";
11
+ import type { ThreadMessage } from "@assistant-ui/core";
6
12
  import type { A2AClient } from "./A2AClient";
7
13
  import type { A2AStreamEvent } from "./types";
8
14
  import { useA2ARuntime } from "./useA2ARuntime";
@@ -68,7 +74,21 @@ const createFetchMock = () =>
68
74
  : input.url;
69
75
  if (url.endsWith("/.well-known/agent-card.json")) {
70
76
  return new Response(
71
- JSON.stringify({ capabilities: { streaming: true } }),
77
+ JSON.stringify({
78
+ name: "Test Agent",
79
+ version: "1.0",
80
+ supported_interfaces: [
81
+ {
82
+ url: "https://agent.test",
83
+ protocol_binding: "HTTP+JSON",
84
+ protocol_version: "1.0",
85
+ },
86
+ ],
87
+ capabilities: { streaming: true },
88
+ default_input_modes: ["text"],
89
+ default_output_modes: ["text"],
90
+ skills: [],
91
+ }),
72
92
  {
73
93
  status: 200,
74
94
  headers: { "Content-Type": "application/json" },
@@ -84,6 +104,15 @@ const createFetchMock = () =>
84
104
  );
85
105
  });
86
106
 
107
+ const createThreadMessage = (id: string): ThreadMessage => ({
108
+ id,
109
+ role: "user",
110
+ content: [{ type: "text", text: id }],
111
+ attachments: [],
112
+ createdAt: new Date(0),
113
+ metadata: { custom: {} },
114
+ });
115
+
87
116
  afterEach(() => {
88
117
  vi.unstubAllGlobals();
89
118
  });
@@ -146,10 +175,12 @@ describe("useA2ARuntime", () => {
146
175
  rerender({ client: second.client });
147
176
 
148
177
  await waitFor(() => expect(second.getAgentCard).toHaveBeenCalledOnce());
149
- await waitFor(() => expect(history.load).toHaveBeenCalledTimes(2));
150
- expect(result.current.thread.getState().messages.map((m) => m.id)).toEqual([
151
- "restored",
152
- ]);
178
+ await waitFor(() =>
179
+ expect(
180
+ result.current.thread.getState().messages.map((m) => m.id),
181
+ ).toEqual(["restored"]),
182
+ );
183
+ expect(history.load).toHaveBeenCalledTimes(2);
153
184
  });
154
185
 
155
186
  it("switches provided clients and aborts the previous client run", async () => {
@@ -277,4 +308,92 @@ describe("useA2ARuntime", () => {
277
308
  Authorization: "Bearer workspace-a",
278
309
  });
279
310
  });
311
+
312
+ it("ignores an older thread load after a newer selection", async () => {
313
+ const { client } = createMockClient();
314
+ let resolveFirst!: (value: { messages: ThreadMessage[] }) => void;
315
+ let resolveSecond!: (value: { messages: ThreadMessage[] }) => void;
316
+ const first = new Promise<{ messages: ThreadMessage[] }>((resolve) => {
317
+ resolveFirst = resolve;
318
+ });
319
+ const second = new Promise<{ messages: ThreadMessage[] }>((resolve) => {
320
+ resolveSecond = resolve;
321
+ });
322
+ const { result } = renderHook(() => {
323
+ const [threadId, setThreadId] = useState("initial");
324
+ return useA2ARuntime({
325
+ client,
326
+ adapters: {
327
+ threadList: {
328
+ threadId,
329
+ onSwitchToThread: async (nextThreadId) => {
330
+ setThreadId(nextThreadId);
331
+ return nextThreadId === "thread-a" ? first : second;
332
+ },
333
+ },
334
+ },
335
+ });
336
+ });
337
+
338
+ let switchA!: Promise<void>;
339
+ let switchB!: Promise<void>;
340
+ act(() => {
341
+ switchA = result.current.threads.switchToThread("thread-a");
342
+ switchB = result.current.threads.switchToThread("thread-b");
343
+ });
344
+ expect(result.current.threads.getState().mainThreadId).toBe("thread-b");
345
+
346
+ await act(async () => {
347
+ resolveSecond({ messages: [createThreadMessage("thread-b")] });
348
+ await switchB;
349
+ });
350
+ await act(async () => {
351
+ resolveFirst({ messages: [createThreadMessage("thread-a")] });
352
+ await switchA;
353
+ });
354
+
355
+ expect(result.current.thread.getState().messages.map((m) => m.id)).toEqual([
356
+ "thread-b",
357
+ ]);
358
+ });
359
+
360
+ it("ignores a thread load superseded by a new thread", async () => {
361
+ const { client } = createMockClient();
362
+ let resolveLoad!: (value: { messages: ThreadMessage[] }) => void;
363
+ const load = new Promise<{ messages: ThreadMessage[] }>((resolve) => {
364
+ resolveLoad = resolve;
365
+ });
366
+ const { result } = renderHook(() => {
367
+ const [threadId, setThreadId] = useState("initial");
368
+ return useA2ARuntime({
369
+ client,
370
+ adapters: {
371
+ threadList: {
372
+ threadId,
373
+ onSwitchToThread: async (nextThreadId) => {
374
+ setThreadId(nextThreadId);
375
+ return load;
376
+ },
377
+ onSwitchToNewThread: async () => {
378
+ setThreadId("thread-new");
379
+ },
380
+ },
381
+ },
382
+ });
383
+ });
384
+
385
+ let staleSwitch!: Promise<void>;
386
+ act(() => {
387
+ staleSwitch = result.current.threads.switchToThread("thread-a");
388
+ });
389
+ await act(async () => {
390
+ await result.current.threads.switchToNewThread();
391
+ });
392
+ await act(async () => {
393
+ resolveLoad({ messages: [createThreadMessage("thread-a")] });
394
+ await staleSwitch;
395
+ });
396
+
397
+ expect(result.current.thread.getState().messages).toEqual([]);
398
+ });
280
399
  });
@@ -110,6 +110,7 @@ export function useA2ARuntime(options: UseA2ARuntimeOptions): AssistantRuntime {
110
110
  });
111
111
 
112
112
  // Thread list
113
+ const threadSwitchGenerationRef = useRef(0);
113
114
  const threadList = useMemo(() => {
114
115
  if (!threadListAdapter) return undefined;
115
116
 
@@ -119,7 +120,9 @@ export function useA2ARuntime(options: UseA2ARuntimeOptions): AssistantRuntime {
119
120
  threadId: threadListAdapter.threadId,
120
121
  onSwitchToNewThread: onSwitchToNewThread
121
122
  ? async () => {
123
+ const generation = ++threadSwitchGenerationRef.current;
122
124
  await onSwitchToNewThread();
125
+ if (generation !== threadSwitchGenerationRef.current) return;
123
126
  // Apply first so the abort inside resetContext finds an already
124
127
  // cleared repository and cannot persist the old thread's partial
125
128
  // assistant message.
@@ -129,7 +132,9 @@ export function useA2ARuntime(options: UseA2ARuntimeOptions): AssistantRuntime {
129
132
  : undefined,
130
133
  onSwitchToThread: onSwitchToThread
131
134
  ? async (threadId: string) => {
135
+ const generation = ++threadSwitchGenerationRef.current;
132
136
  const result = await onSwitchToThread(threadId);
137
+ if (generation !== threadSwitchGenerationRef.current) return;
133
138
  core.applyExternalMessages(result.messages);
134
139
  core.resetContext();
135
140
  }