@assistant-ui/react-a2a 0.2.31 → 0.2.32

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.
@@ -10,13 +10,15 @@ import type {
10
10
  ThreadHistoryAdapter,
11
11
  ThreadMessage,
12
12
  } from "@assistant-ui/core";
13
- import { MessageRepository } from "@assistant-ui/core/internal";
13
+ import {
14
+ createMessageRepositorySession,
15
+ invokeUserCallback,
16
+ } from "@assistant-ui/core/internal";
14
17
  import type { A2AClient } from "./A2AClient";
15
18
  import type {
16
19
  A2AArtifact,
17
20
  A2AAgentCard,
18
21
  A2AMessage,
19
- A2APart,
20
22
  A2ASendMessageConfiguration,
21
23
  A2AStreamEvent,
22
24
  A2ATask,
@@ -25,8 +27,8 @@ import type {
25
27
  } from "./types";
26
28
  import {
27
29
  a2aMessageToContent,
28
- contentPartsToA2AParts,
29
30
  isTerminalTaskState,
31
+ threadMessageToA2AMessage,
30
32
  taskStateToMessageStatus,
31
33
  } from "./conversions";
32
34
 
@@ -48,32 +50,12 @@ const FALLBACK_USER_STATUS = {
48
50
 
49
51
  type A2ARuntimeCallbackName = "onError" | "onCancel" | "onArtifactComplete";
50
52
 
51
- const reportCallbackError = (name: A2ARuntimeCallbackName, error: unknown) => {
52
- console.error(`[react-a2a] ${name} callback threw an error`, error);
53
- };
54
-
55
- const invokeRuntimeCallback = <TArgs extends unknown[]>(
53
+ const invokeRuntimeCallback = <TArgs extends readonly unknown[]>(
56
54
  name: A2ARuntimeCallbackName,
57
- callback: ((...args: TArgs) => void) | undefined,
55
+ callback: ((...args: TArgs) => unknown) | undefined,
58
56
  ...args: TArgs
59
- ) => {
60
- if (!callback) return;
61
-
62
- try {
63
- const result = callback(...args) as unknown;
64
- if (
65
- result !== null &&
66
- (typeof result === "object" || typeof result === "function") &&
67
- "then" in result &&
68
- typeof result.then === "function"
69
- ) {
70
- void Promise.resolve(result).catch((error) => {
71
- reportCallbackError(name, error);
72
- });
73
- }
74
- } catch (error) {
75
- reportCallbackError(name, error);
76
- }
57
+ ): void => {
58
+ void invokeUserCallback("react-a2a", name, callback, ...args);
77
59
  };
78
60
 
79
61
  function normalizeArtifact(artifact: A2AArtifact): A2AArtifact {
@@ -94,8 +76,7 @@ export class A2AThreadRuntimeCore {
94
76
  private readonly notifyUpdate: () => void;
95
77
 
96
78
  private runtime: AssistantRuntime | undefined;
97
- private readonly repository = new MessageRepository();
98
- private exportedRepository: ExportedMessageRepository | undefined;
79
+ private readonly session = createMessageRepositorySession();
99
80
  private isRunningFlag = false;
100
81
  private abortController: AbortController | null = null;
101
82
  private pendingError: Error | null = null;
@@ -110,10 +91,15 @@ export class A2AThreadRuntimeCore {
110
91
  private readonly recordedHistoryIds = new Set<string>();
111
92
  private _isLoading = false;
112
93
  private _loadPromise: Promise<void> | undefined;
94
+ private _loadRequested = false;
95
+ private _agentCardPromise: Promise<void> | undefined;
96
+
97
+ private lastOptionsContextId: string | undefined;
113
98
 
114
99
  constructor(options: A2AThreadRuntimeCoreOptions) {
115
100
  this.client = options.client;
116
101
  this.contextId = options.contextId;
102
+ this.lastOptionsContextId = options.contextId;
117
103
  this.configuration = options.configuration;
118
104
  this.onError = options.onError;
119
105
  this.onCancel = options.onCancel;
@@ -124,12 +110,46 @@ export class A2AThreadRuntimeCore {
124
110
 
125
111
  updateOptions(options: Omit<A2AThreadRuntimeCoreOptions, "notifyUpdate">) {
126
112
  this.client = options.client;
127
- this.contextId = options.contextId;
113
+ // The hook re-applies options on every render, including renders caused
114
+ // by this core's own notifyUpdate. The option only seeds the context: a
115
+ // re-render with the same value must not clobber a server-assigned
116
+ // contextId learned from the stream.
117
+ if (options.contextId !== this.lastOptionsContextId) {
118
+ this.contextId = options.contextId;
119
+ this.lastOptionsContextId = options.contextId;
120
+ }
128
121
  this.configuration = options.configuration;
129
122
  this.onError = options.onError;
130
123
  this.onCancel = options.onCancel;
131
124
  this.onArtifactComplete = options.onArtifactComplete;
125
+ const previousHistory = this.history;
132
126
  this.history = options.history;
127
+
128
+ if (
129
+ this._loadRequested &&
130
+ !this._loadPromise &&
131
+ !previousHistory &&
132
+ options.history &&
133
+ this.session.getMessages().length === 0
134
+ ) {
135
+ void this.__internal_load();
136
+ }
137
+ }
138
+
139
+ /** Thread-boundary reset: applyExternalMessages alone also serves branch
140
+ * switches, deletes, and cancel resyncs, which must keep the live context. */
141
+ resetContext(): void {
142
+ // Restore the seed before aborting: an onCancel callback that starts a
143
+ // new run must not pick up the old thread's context, and its controller
144
+ // must not be discarded.
145
+ const controller = this.abortController;
146
+ this.contextId = this.lastOptionsContextId;
147
+ if (controller) {
148
+ controller.abort();
149
+ if (this.abortController === controller) {
150
+ this.abortController = null;
151
+ }
152
+ }
133
153
  }
134
154
 
135
155
  attachRuntime(runtime: AssistantRuntime) {
@@ -150,69 +170,11 @@ export class A2AThreadRuntimeCore {
150
170
  }
151
171
 
152
172
  getMessages(): readonly ThreadMessage[] {
153
- return this.repository.getMessages();
173
+ return this.session.getMessages();
154
174
  }
155
175
 
156
176
  getMessageRepository(): ExportedMessageRepository {
157
- this.exportedRepository ??= this.repository.export();
158
- return this.exportedRepository;
159
- }
160
-
161
- private tryGetMessage(messageId: string) {
162
- try {
163
- return this.repository.getMessage(messageId);
164
- } catch {
165
- return undefined;
166
- }
167
- }
168
-
169
- private tryGetMessages(
170
- messageId: string,
171
- ): readonly ThreadMessage[] | undefined {
172
- try {
173
- return this.repository.getMessages(messageId);
174
- } catch {
175
- return undefined;
176
- }
177
- }
178
-
179
- private hasMessage(messageId: string): boolean {
180
- return this.tryGetMessage(messageId) !== undefined;
181
- }
182
-
183
- private addOrUpdateMessage(
184
- parentId: string | null,
185
- message: ThreadMessage,
186
- ): void {
187
- this.repository.addOrUpdateMessage(parentId, message);
188
- this.exportedRepository = undefined;
189
- }
190
-
191
- private switchToBranch(messageId: string): void {
192
- this.repository.switchToBranch(messageId);
193
- this.exportedRepository = undefined;
194
- }
195
-
196
- private resetRepositoryHead(messageId: string | null): void {
197
- this.repository.resetHead(messageId);
198
- this.exportedRepository = undefined;
199
- }
200
-
201
- private clearRepository(): void {
202
- this.repository.clear();
203
- this.exportedRepository = undefined;
204
- }
205
-
206
- private updateMessage(
207
- messageId: string,
208
- updater: (message: ThreadMessage) => ThreadMessage,
209
- ): boolean {
210
- const item = this.tryGetMessage(messageId);
211
- if (!item) return false;
212
- const message = updater(item.message);
213
- if (message === item.message) return false;
214
- this.addOrUpdateMessage(item.parentId, message);
215
- return true;
177
+ return this.session.export();
216
178
  }
217
179
 
218
180
  getTask(): A2ATask | undefined {
@@ -236,20 +198,27 @@ export class A2AThreadRuntimeCore {
236
198
  }
237
199
 
238
200
  __internal_load(): Promise<void> {
201
+ this._loadRequested = true;
202
+ this._agentCardPromise ??= this.client
203
+ .getAgentCard()
204
+ .then((agentCard) => {
205
+ this.agentCardValue = agentCard;
206
+ this.notifyUpdate();
207
+ })
208
+ .catch(() => undefined);
209
+
239
210
  if (this._loadPromise) return this._loadPromise;
211
+ if (!this.history) return this._agentCardPromise;
240
212
 
241
213
  this._isLoading = true;
242
214
 
243
- const historyPromise = this.history?.load() ?? Promise.resolve(null);
244
- const agentCardPromise = this.client.getAgentCard().catch(() => undefined);
215
+ const historyPromise = this.history.load();
245
216
 
246
- this._loadPromise = Promise.all([historyPromise, agentCardPromise])
247
- .then(([repo, agentCard]) => {
248
- if (agentCard) {
249
- this.agentCardValue = agentCard;
250
- }
217
+ this._loadPromise = Promise.all([historyPromise, this._agentCardPromise])
218
+ .then(([repo]) => {
251
219
  if (repo) {
252
- this.applyExternalMessageRepository(repo);
220
+ this.session.applyExternalMessageRepository(repo);
221
+ this.finalizeExternalApply();
253
222
  }
254
223
  })
255
224
  .catch((error) => {
@@ -279,11 +248,11 @@ export class A2AThreadRuntimeCore {
279
248
  const parentId =
280
249
  message.parentId === null
281
250
  ? null
282
- : message.parentId && this.hasMessage(message.parentId)
251
+ : message.parentId && this.session.hasMessage(message.parentId)
283
252
  ? message.parentId
284
- : this.repository.headId;
285
- this.addOrUpdateMessage(parentId, threadMessage);
286
- this.switchToBranch(threadMessage.id);
253
+ : this.session.headId;
254
+ this.session.addOrUpdateMessage(parentId, threadMessage);
255
+ this.session.switchToBranch(threadMessage.id);
287
256
  this.notifyUpdate();
288
257
  this.recordHistoryEntry(parentId, threadMessage);
289
258
 
@@ -302,7 +271,7 @@ export class A2AThreadRuntimeCore {
302
271
  const messages =
303
272
  parentId === null
304
273
  ? []
305
- : (this.tryGetMessages(parentId) ?? this.getMessages());
274
+ : (this.session.tryGetMessages(parentId) ?? this.getMessages());
306
275
  for (let i = messages.length - 1; i >= 0; i--) {
307
276
  if (messages[i]!.role === "user") {
308
277
  await this.startRun(messages[i]!);
@@ -336,7 +305,7 @@ export class A2AThreadRuntimeCore {
336
305
  for (const message of messages) {
337
306
  if (seen.has(message.id)) continue;
338
307
  seen.add(message.id);
339
- this.addOrUpdateMessage(parentId, message);
308
+ this.session.addOrUpdateMessage(parentId, message);
340
309
  parentId = message.id;
341
310
  lastId = message.id;
342
311
  }
@@ -357,7 +326,7 @@ export class A2AThreadRuntimeCore {
357
326
 
358
327
  applyExternalMessages(messages: readonly ThreadMessage[]): void {
359
328
  if (messages.length === 0) {
360
- this.clearRepository();
329
+ this.session.clear();
361
330
  } else {
362
331
  let expectedParentId: string | null = null;
363
332
  let lastAppliedId: string | null = null;
@@ -367,81 +336,22 @@ export class A2AThreadRuntimeCore {
367
336
  for (const message of messages) {
368
337
  if (seen.has(message.id)) continue;
369
338
  seen.add(message.id);
370
- const existing = this.tryGetMessage(message.id);
339
+ const existing = this.session.tryGetMessage(message.id);
371
340
  if (existing && existing.parentId !== expectedParentId) {
372
341
  hardReplace = true;
373
342
  break;
374
343
  }
375
- this.addOrUpdateMessage(expectedParentId, message);
344
+ this.session.addOrUpdateMessage(expectedParentId, message);
376
345
  expectedParentId = message.id;
377
346
  lastAppliedId = message.id;
378
347
  }
379
348
 
380
349
  if (hardReplace) {
381
- this.clearRepository();
350
+ this.session.clear();
382
351
  lastAppliedId = this.appendLinearChain(messages);
383
352
  }
384
353
 
385
- this.resetRepositoryHead(lastAppliedId);
386
- }
387
-
388
- this.finalizeExternalApply();
389
- }
390
-
391
- private applyExternalMessageRepository(
392
- loaded: ExportedMessageRepository,
393
- ): void {
394
- const headId = loaded.headId ?? loaded.messages.at(-1)?.message.id ?? null;
395
- const ids = new Set<string>();
396
- let degenerate = false;
397
- for (const { message } of loaded.messages) {
398
- if (ids.has(message.id)) {
399
- degenerate = true;
400
- break;
401
- }
402
- ids.add(message.id);
403
- }
404
- if (headId !== null && !ids.has(headId)) degenerate = true;
405
-
406
- if (!degenerate) {
407
- this.clearRepository();
408
- let pending = [...loaded.messages];
409
- const importedIds = new Set<string>();
410
-
411
- while (pending.length > 0) {
412
- const unresolved: typeof pending = [];
413
- let progressed = false;
414
- for (const item of pending) {
415
- if (item.parentId !== null && !importedIds.has(item.parentId)) {
416
- unresolved.push(item);
417
- continue;
418
- }
419
- this.addOrUpdateMessage(item.parentId, item.message);
420
- importedIds.add(item.message.id);
421
- progressed = true;
422
- }
423
- if (!progressed) {
424
- degenerate = true;
425
- break;
426
- }
427
- pending = unresolved;
428
- }
429
- }
430
-
431
- if (degenerate) {
432
- this.clearRepository();
433
- let previousId: string | null = null;
434
- for (const { message } of loaded.messages) {
435
- const existing = this.tryGetMessage(message.id);
436
- this.addOrUpdateMessage(
437
- existing ? existing.parentId : previousId,
438
- message,
439
- );
440
- previousId = message.id;
441
- }
442
- this.resetRepositoryHead(previousId);
443
- } else {
444
- this.resetRepositoryHead(headId);
354
+ this.session.resetHead(lastAppliedId);
445
355
  }
446
356
 
447
357
  this.finalizeExternalApply();
@@ -456,7 +366,14 @@ export class A2AThreadRuntimeCore {
456
366
  this.abortController = null;
457
367
  }
458
368
 
459
- const a2aMessage = this.threadMessageToA2AMessage(userThreadMessage);
369
+ const a2aMessage = threadMessageToA2AMessage(userThreadMessage, {
370
+ contextId: this.contextId,
371
+ taskId:
372
+ this.currentTask?.id &&
373
+ !isTerminalTaskState(this.currentTask.status.state)
374
+ ? this.currentTask.id
375
+ : undefined,
376
+ });
460
377
 
461
378
  // Clear task if previous task reached terminal state
462
379
  if (
@@ -728,41 +645,6 @@ export class A2AThreadRuntimeCore {
728
645
 
729
646
  // --- Message helpers ---
730
647
 
731
- private threadMessageToA2AMessage(message: ThreadMessage): A2AMessage {
732
- const parts: A2APart[] = [];
733
-
734
- if (message.role === "user") {
735
- parts.push(...contentPartsToA2AParts(message.content));
736
- for (const attachment of message.attachments ?? []) {
737
- parts.push(
738
- ...contentPartsToA2AParts(
739
- attachment.content ?? [],
740
- attachment.contentType,
741
- ),
742
- );
743
- }
744
- }
745
-
746
- const a2aMsg: A2AMessage = {
747
- messageId: message.id,
748
- role: "user",
749
- parts,
750
- };
751
-
752
- if (this.contextId) {
753
- a2aMsg.contextId = this.contextId;
754
- }
755
- // Only attach taskId if current task is NOT in terminal state
756
- if (
757
- this.currentTask?.id &&
758
- !isTerminalTaskState(this.currentTask.status.state)
759
- ) {
760
- a2aMsg.taskId = this.currentTask.id;
761
- }
762
-
763
- return a2aMsg;
764
- }
765
-
766
648
  private insertAssistantPlaceholder(parentId: string): string {
767
649
  const id = generateId();
768
650
  const assistant: ThreadAssistantMessage = {
@@ -779,8 +661,8 @@ export class A2AThreadRuntimeCore {
779
661
  custom: {},
780
662
  },
781
663
  };
782
- this.addOrUpdateMessage(parentId, assistant);
783
- this.switchToBranch(id);
664
+ this.session.addOrUpdateMessage(parentId, assistant);
665
+ this.session.switchToBranch(id);
784
666
  this.notifyUpdate();
785
667
  return id;
786
668
  }
@@ -789,14 +671,14 @@ export class A2AThreadRuntimeCore {
789
671
  messageId: string,
790
672
  content: ThreadAssistantMessage["content"],
791
673
  ) {
792
- this.updateMessage(messageId, (message) => {
674
+ this.session.updateMessage(messageId, (message) => {
793
675
  if (message.role !== "assistant") return message;
794
676
  return { ...message, content };
795
677
  });
796
678
  }
797
679
 
798
680
  private updateAssistantStatus(messageId: string, status: MessageStatus) {
799
- const touched = this.updateMessage(messageId, (message) => {
681
+ const touched = this.session.updateMessage(messageId, (message) => {
800
682
  if (message.role !== "assistant") return message;
801
683
  return { ...message, status };
802
684
  });
@@ -809,7 +691,7 @@ export class A2AThreadRuntimeCore {
809
691
  }
810
692
 
811
693
  private getAssistantStatus(messageId: string): MessageStatus | undefined {
812
- const msg = this.tryGetMessage(messageId)?.message;
694
+ const msg = this.session.tryGetMessage(messageId)?.message;
813
695
  if (msg?.role !== "assistant") return undefined;
814
696
  return msg.status;
815
697
  }
@@ -845,7 +727,7 @@ export class A2AThreadRuntimeCore {
845
727
  if (!this.history) return;
846
728
  const parentId = this.assistantHistoryParents.get(messageId);
847
729
  if (parentId === undefined) return;
848
- const message = this.tryGetMessage(messageId)?.message;
730
+ const message = this.session.tryGetMessage(messageId)?.message;
849
731
  if (!message || message.role !== "assistant") return;
850
732
  if (
851
733
  message.status?.type !== "complete" &&
@@ -7,6 +7,7 @@ import {
7
7
  contentPartsToA2AParts,
8
8
  isTerminalTaskState,
9
9
  isInterruptedTaskState,
10
+ threadMessageToA2AMessage,
10
11
  } from "./conversions";
11
12
  import type { A2APart, A2AMessage, A2ATaskState } from "./types";
12
13
 
@@ -628,3 +629,51 @@ describe("contentPartsToA2AParts", () => {
628
629
  expect(contentPartsToA2AParts([])).toEqual([]);
629
630
  });
630
631
  });
632
+
633
+ describe("threadMessageToA2AMessage", () => {
634
+ const userMessage = {
635
+ id: "msg-1",
636
+ role: "user",
637
+ createdAt: new Date(),
638
+ content: [{ type: "text" as const, text: "hello" }],
639
+ attachments: [
640
+ {
641
+ id: "att-1",
642
+ type: "file" as const,
643
+ name: "notes.txt",
644
+ contentType: "text/plain",
645
+ status: { type: "complete" as const },
646
+ content: [{ type: "text" as const, text: "attached" }],
647
+ },
648
+ ],
649
+ metadata: { custom: {} },
650
+ } as any;
651
+
652
+ it("converts user content and appends attachment parts", () => {
653
+ const result = threadMessageToA2AMessage(userMessage);
654
+ expect(result.messageId).toBe("msg-1");
655
+ expect(result.role).toBe("user");
656
+ expect(result.parts).toEqual([{ text: "hello" }, { text: "attached" }]);
657
+ expect(result.contextId).toBeUndefined();
658
+ expect(result.taskId).toBeUndefined();
659
+ });
660
+
661
+ it("attaches contextId and taskId when provided", () => {
662
+ const result = threadMessageToA2AMessage(userMessage, {
663
+ contextId: "ctx-1",
664
+ taskId: "task-1",
665
+ });
666
+ expect(result.contextId).toBe("ctx-1");
667
+ expect(result.taskId).toBe("task-1");
668
+ });
669
+
670
+ it("skips undefined options and non-user content", () => {
671
+ const result = threadMessageToA2AMessage(
672
+ { ...userMessage, role: "assistant" },
673
+ { contextId: undefined, taskId: undefined },
674
+ );
675
+ expect(result.parts).toEqual([]);
676
+ expect(result.contextId).toBeUndefined();
677
+ expect(result.taskId).toBeUndefined();
678
+ });
679
+ });
@@ -1,7 +1,14 @@
1
1
  "use client";
2
2
 
3
- import type { MessageStatus, ThreadAssistantMessage } from "@assistant-ui/core";
4
- import { httpUrlPattern, parseDataUrl } from "@assistant-ui/core/internal";
3
+ import type {
4
+ MessageStatus,
5
+ ThreadAssistantMessage,
6
+ ThreadMessage,
7
+ } from "@assistant-ui/core";
8
+ import {
9
+ parseDataUrl,
10
+ resolveFilePartSource,
11
+ } from "@assistant-ui/core/internal";
5
12
  import type { A2AMessage, A2APart, A2ATaskState } from "./types";
6
13
 
7
14
  function isImageMediaType(mediaType?: string): boolean {
@@ -134,9 +141,14 @@ export function contentPartsToA2AParts(
134
141
  case "file": {
135
142
  if (typeof part.data !== "string" || !part.data) return null;
136
143
  const declaredMimeType = part.mimeType || fallbackMimeType;
137
- if (part.sourceType === "url" || httpUrlPattern.test(part.data)) {
144
+ const source = resolveFilePartSource({
145
+ data: part.data,
146
+ mimeType: declaredMimeType ?? "application/octet-stream",
147
+ sourceType: part.sourceType,
148
+ });
149
+ if (source.kind === "url") {
138
150
  return {
139
- url: part.data,
151
+ url: source.url,
140
152
  ...(declaredMimeType && { mediaType: declaredMimeType }),
141
153
  ...(part.filename && { filename: part.filename }),
142
154
  };
@@ -144,8 +156,8 @@ export function contentPartsToA2AParts(
144
156
  const parsed = parseDataUrl(part.data);
145
157
  if (parsed) {
146
158
  return {
147
- raw: parsed.data,
148
- mediaType: parsed.mimeType,
159
+ raw: source.data,
160
+ mediaType: source.mimeType,
149
161
  ...(part.filename && { filename: part.filename }),
150
162
  };
151
163
  }
@@ -157,7 +169,7 @@ export function contentPartsToA2AParts(
157
169
  };
158
170
  }
159
171
  return {
160
- raw: part.data,
172
+ raw: source.data,
161
173
  ...(declaredMimeType && { mediaType: declaredMimeType }),
162
174
  ...(part.filename && { filename: part.filename }),
163
175
  };
@@ -185,3 +197,40 @@ export function a2aMessageToContent(
185
197
  ): ThreadAssistantMessage["content"] {
186
198
  return a2aPartsToContent(message?.parts ?? []);
187
199
  }
200
+
201
+ export function threadMessageToA2AMessage(
202
+ message: ThreadMessage,
203
+ options: {
204
+ contextId?: string | undefined;
205
+ taskId?: string | undefined;
206
+ } = {},
207
+ ): A2AMessage {
208
+ const parts: A2APart[] = [];
209
+
210
+ if (message.role === "user") {
211
+ parts.push(...contentPartsToA2AParts(message.content));
212
+ for (const attachment of message.attachments ?? []) {
213
+ parts.push(
214
+ ...contentPartsToA2AParts(
215
+ attachment.content ?? [],
216
+ attachment.contentType,
217
+ ),
218
+ );
219
+ }
220
+ }
221
+
222
+ const a2aMsg: A2AMessage = {
223
+ messageId: message.id,
224
+ role: "user",
225
+ parts,
226
+ };
227
+
228
+ if (options.contextId) {
229
+ a2aMsg.contextId = options.contextId;
230
+ }
231
+ if (options.taskId) {
232
+ a2aMsg.taskId = options.taskId;
233
+ }
234
+
235
+ return a2aMsg;
236
+ }