@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.
@@ -4,11 +4,13 @@ import { generateId, fromThreadMessageLike } from "@assistant-ui/core";
4
4
  import type {
5
5
  AppendMessage,
6
6
  AssistantRuntime,
7
+ ExportedMessageRepository,
7
8
  MessageStatus,
8
9
  ThreadAssistantMessage,
9
10
  ThreadHistoryAdapter,
10
11
  ThreadMessage,
11
12
  } from "@assistant-ui/core";
13
+ import { MessageRepository } from "@assistant-ui/core/internal";
12
14
  import type { A2AClient } from "./A2AClient";
13
15
  import type {
14
16
  A2AArtifact,
@@ -23,6 +25,7 @@ import type {
23
25
  } from "./types";
24
26
  import {
25
27
  a2aMessageToContent,
28
+ contentPartsToA2AParts,
26
29
  isTerminalTaskState,
27
30
  taskStateToMessageStatus,
28
31
  } from "./conversions";
@@ -54,7 +57,8 @@ export class A2AThreadRuntimeCore {
54
57
  private readonly notifyUpdate: () => void;
55
58
 
56
59
  private runtime: AssistantRuntime | undefined;
57
- private messages: ThreadMessage[] = [];
60
+ private readonly repository = new MessageRepository();
61
+ private exportedRepository: ExportedMessageRepository | undefined;
58
62
  private isRunningFlag = false;
59
63
  private abortController: AbortController | null = null;
60
64
  private pendingError: Error | null = null;
@@ -109,7 +113,69 @@ export class A2AThreadRuntimeCore {
109
113
  }
110
114
 
111
115
  getMessages(): readonly ThreadMessage[] {
112
- return this.messages;
116
+ return this.repository.getMessages();
117
+ }
118
+
119
+ getMessageRepository(): ExportedMessageRepository {
120
+ this.exportedRepository ??= this.repository.export();
121
+ return this.exportedRepository;
122
+ }
123
+
124
+ private tryGetMessage(messageId: string) {
125
+ try {
126
+ return this.repository.getMessage(messageId);
127
+ } catch {
128
+ return undefined;
129
+ }
130
+ }
131
+
132
+ private tryGetMessages(
133
+ messageId: string,
134
+ ): readonly ThreadMessage[] | undefined {
135
+ try {
136
+ return this.repository.getMessages(messageId);
137
+ } catch {
138
+ return undefined;
139
+ }
140
+ }
141
+
142
+ private hasMessage(messageId: string): boolean {
143
+ return this.tryGetMessage(messageId) !== undefined;
144
+ }
145
+
146
+ private addOrUpdateMessage(
147
+ parentId: string | null,
148
+ message: ThreadMessage,
149
+ ): void {
150
+ this.repository.addOrUpdateMessage(parentId, message);
151
+ this.exportedRepository = undefined;
152
+ }
153
+
154
+ private switchToBranch(messageId: string): void {
155
+ this.repository.switchToBranch(messageId);
156
+ this.exportedRepository = undefined;
157
+ }
158
+
159
+ private resetRepositoryHead(messageId: string | null): void {
160
+ this.repository.resetHead(messageId);
161
+ this.exportedRepository = undefined;
162
+ }
163
+
164
+ private clearRepository(): void {
165
+ this.repository.clear();
166
+ this.exportedRepository = undefined;
167
+ }
168
+
169
+ private updateMessage(
170
+ messageId: string,
171
+ updater: (message: ThreadMessage) => ThreadMessage,
172
+ ): boolean {
173
+ const item = this.tryGetMessage(messageId);
174
+ if (!item) return false;
175
+ const message = updater(item.message);
176
+ if (message === item.message) return false;
177
+ this.addOrUpdateMessage(item.parentId, message);
178
+ return true;
113
179
  }
114
180
 
115
181
  getTask(): A2ATask | undefined {
@@ -146,8 +212,7 @@ export class A2AThreadRuntimeCore {
146
212
  this.agentCardValue = agentCard;
147
213
  }
148
214
  if (repo) {
149
- const messages = repo.messages.map((item) => item.message);
150
- this.applyExternalMessages(messages);
215
+ this.applyExternalMessageRepository(repo);
151
216
  }
152
217
  })
153
218
  .catch((error) => {
@@ -166,21 +231,22 @@ export class A2AThreadRuntimeCore {
166
231
 
167
232
  async append(message: AppendMessage): Promise<void> {
168
233
  const startRun = message.startRun ?? message.role === "user";
169
- if (message.sourceId) {
170
- this.messages = this.messages.filter(
171
- (entry) => entry.id !== message.sourceId,
172
- );
173
- }
174
- this.resetHead(message.parentId);
175
234
 
176
235
  const threadMessage = fromThreadMessageLike(
177
236
  message as any,
178
237
  generateId(),
179
238
  FALLBACK_USER_STATUS,
180
239
  );
181
- this.messages = [...this.messages, threadMessage];
240
+ const parentId =
241
+ message.parentId === null
242
+ ? null
243
+ : message.parentId && this.hasMessage(message.parentId)
244
+ ? message.parentId
245
+ : this.repository.headId;
246
+ this.addOrUpdateMessage(parentId, threadMessage);
247
+ this.switchToBranch(threadMessage.id);
182
248
  this.notifyUpdate();
183
- this.recordHistoryEntry(message.parentId ?? null, threadMessage);
249
+ this.recordHistoryEntry(parentId, threadMessage);
184
250
 
185
251
  if (!startRun) return;
186
252
  await this.startRun(threadMessage);
@@ -194,13 +260,13 @@ export class A2AThreadRuntimeCore {
194
260
  parentId: string | null,
195
261
  _config: { runConfig?: Record<string, unknown> } = {},
196
262
  ): Promise<void> {
197
- this.resetHead(parentId);
198
- this.notifyUpdate();
199
-
200
- // Find the last user message to re-run
201
- for (let i = this.messages.length - 1; i >= 0; i--) {
202
- if (this.messages[i]!.role === "user") {
203
- await this.startRun(this.messages[i]!);
263
+ const messages =
264
+ parentId === null
265
+ ? []
266
+ : (this.tryGetMessages(parentId) ?? this.getMessages());
267
+ for (let i = messages.length - 1; i >= 0; i--) {
268
+ if (messages[i]!.role === "user") {
269
+ await this.startRun(messages[i]!);
204
270
  return;
205
271
  }
206
272
  }
@@ -223,19 +289,125 @@ export class A2AThreadRuntimeCore {
223
289
  }
224
290
  }
225
291
 
226
- applyExternalMessages(messages: readonly ThreadMessage[]): void {
292
+ private appendLinearChain(messages: readonly ThreadMessage[]): string | null {
293
+ let parentId: string | null = null;
294
+ let lastId: string | null = null;
295
+ const seen = new Set<string>();
296
+
297
+ for (const message of messages) {
298
+ if (seen.has(message.id)) continue;
299
+ seen.add(message.id);
300
+ this.addOrUpdateMessage(parentId, message);
301
+ parentId = message.id;
302
+ lastId = message.id;
303
+ }
304
+
305
+ return lastId;
306
+ }
307
+
308
+ private finalizeExternalApply(): void {
227
309
  this.assistantHistoryParents.clear();
228
- this.messages = [...messages];
229
310
  this.recordedHistoryIds.clear();
230
- for (const message of this.messages) {
311
+ for (const { message } of this.getMessageRepository().messages) {
231
312
  this.recordedHistoryIds.add(message.id);
232
313
  }
233
- // Reset task-specific state to prevent leaking into new thread
234
314
  this.currentTask = undefined;
235
315
  this.currentArtifacts = [];
236
316
  this.notifyUpdate();
237
317
  }
238
318
 
319
+ applyExternalMessages(messages: readonly ThreadMessage[]): void {
320
+ if (messages.length === 0) {
321
+ this.clearRepository();
322
+ } else {
323
+ let expectedParentId: string | null = null;
324
+ let lastAppliedId: string | null = null;
325
+ let hardReplace = false;
326
+ const seen = new Set<string>();
327
+
328
+ for (const message of messages) {
329
+ if (seen.has(message.id)) continue;
330
+ seen.add(message.id);
331
+ const existing = this.tryGetMessage(message.id);
332
+ if (existing && existing.parentId !== expectedParentId) {
333
+ hardReplace = true;
334
+ break;
335
+ }
336
+ this.addOrUpdateMessage(expectedParentId, message);
337
+ expectedParentId = message.id;
338
+ lastAppliedId = message.id;
339
+ }
340
+
341
+ if (hardReplace) {
342
+ this.clearRepository();
343
+ lastAppliedId = this.appendLinearChain(messages);
344
+ }
345
+
346
+ this.resetRepositoryHead(lastAppliedId);
347
+ }
348
+
349
+ this.finalizeExternalApply();
350
+ }
351
+
352
+ private applyExternalMessageRepository(
353
+ loaded: ExportedMessageRepository,
354
+ ): void {
355
+ const headId = loaded.headId ?? loaded.messages.at(-1)?.message.id ?? null;
356
+ const ids = new Set<string>();
357
+ let degenerate = false;
358
+ for (const { message } of loaded.messages) {
359
+ if (ids.has(message.id)) {
360
+ degenerate = true;
361
+ break;
362
+ }
363
+ ids.add(message.id);
364
+ }
365
+ if (headId !== null && !ids.has(headId)) degenerate = true;
366
+
367
+ if (!degenerate) {
368
+ this.clearRepository();
369
+ let pending = [...loaded.messages];
370
+ const importedIds = new Set<string>();
371
+
372
+ while (pending.length > 0) {
373
+ const unresolved: typeof pending = [];
374
+ let progressed = false;
375
+ for (const item of pending) {
376
+ if (item.parentId !== null && !importedIds.has(item.parentId)) {
377
+ unresolved.push(item);
378
+ continue;
379
+ }
380
+ this.addOrUpdateMessage(item.parentId, item.message);
381
+ importedIds.add(item.message.id);
382
+ progressed = true;
383
+ }
384
+ if (!progressed) {
385
+ degenerate = true;
386
+ break;
387
+ }
388
+ pending = unresolved;
389
+ }
390
+ }
391
+
392
+ if (degenerate) {
393
+ this.clearRepository();
394
+ let previousId: string | null = null;
395
+ for (const { message } of loaded.messages) {
396
+ const existing = this.tryGetMessage(message.id);
397
+ this.addOrUpdateMessage(
398
+ existing ? existing.parentId : previousId,
399
+ message,
400
+ );
401
+ previousId = message.id;
402
+ }
403
+ this.resetRepositoryHead(previousId);
404
+ } else {
405
+ this.resetRepositoryHead(headId);
406
+ }
407
+
408
+ this.finalizeExternalApply();
409
+ }
410
+
239
411
  // --- Run logic ---
240
412
 
241
413
  private async startRun(userThreadMessage: ThreadMessage): Promise<void> {
@@ -258,7 +430,7 @@ export class A2AThreadRuntimeCore {
258
430
  this.currentArtifacts = [];
259
431
 
260
432
  const assistantParentId = userThreadMessage.id;
261
- const assistantId = this.insertAssistantPlaceholder();
433
+ const assistantId = this.insertAssistantPlaceholder(assistantParentId);
262
434
  this.markPendingAssistantHistory(assistantId, assistantParentId);
263
435
 
264
436
  const abortController = new AbortController();
@@ -484,12 +656,14 @@ export class A2AThreadRuntimeCore {
484
656
  const parts: A2APart[] = [];
485
657
 
486
658
  if (message.role === "user") {
487
- for (const part of message.content) {
488
- if (part.type === "text") {
489
- parts.push({ text: part.text });
490
- } else if (part.type === "image") {
491
- parts.push({ url: part.image, mediaType: "image/*" });
492
- }
659
+ parts.push(...contentPartsToA2AParts(message.content));
660
+ for (const attachment of message.attachments ?? []) {
661
+ parts.push(
662
+ ...contentPartsToA2AParts(
663
+ attachment.content ?? [],
664
+ attachment.contentType,
665
+ ),
666
+ );
493
667
  }
494
668
  }
495
669
 
@@ -513,7 +687,7 @@ export class A2AThreadRuntimeCore {
513
687
  return a2aMsg;
514
688
  }
515
689
 
516
- private insertAssistantPlaceholder(): string {
690
+ private insertAssistantPlaceholder(parentId: string): string {
517
691
  const id = generateId();
518
692
  const assistant: ThreadAssistantMessage = {
519
693
  id,
@@ -529,7 +703,8 @@ export class A2AThreadRuntimeCore {
529
703
  custom: {},
530
704
  },
531
705
  };
532
- this.messages = [...this.messages, assistant];
706
+ this.addOrUpdateMessage(parentId, assistant);
707
+ this.switchToBranch(id);
533
708
  this.notifyUpdate();
534
709
  return id;
535
710
  }
@@ -538,19 +713,15 @@ export class A2AThreadRuntimeCore {
538
713
  messageId: string,
539
714
  content: ThreadAssistantMessage["content"],
540
715
  ) {
541
- this.messages = this.messages.map((message) => {
542
- if (message.id !== messageId || message.role !== "assistant")
543
- return message;
716
+ this.updateMessage(messageId, (message) => {
717
+ if (message.role !== "assistant") return message;
544
718
  return { ...message, content };
545
719
  });
546
720
  }
547
721
 
548
722
  private updateAssistantStatus(messageId: string, status: MessageStatus) {
549
- let touched = false;
550
- this.messages = this.messages.map((message) => {
551
- if (message.id !== messageId || message.role !== "assistant")
552
- return message;
553
- touched = true;
723
+ const touched = this.updateMessage(messageId, (message) => {
724
+ if (message.role !== "assistant") return message;
554
725
  return { ...message, status };
555
726
  });
556
727
  if (touched) {
@@ -562,10 +733,9 @@ export class A2AThreadRuntimeCore {
562
733
  }
563
734
 
564
735
  private getAssistantStatus(messageId: string): MessageStatus | undefined {
565
- const msg = this.messages.find(
566
- (m) => m.id === messageId && m.role === "assistant",
567
- );
568
- return msg?.status;
736
+ const msg = this.tryGetMessage(messageId)?.message;
737
+ if (msg?.role !== "assistant") return undefined;
738
+ return msg.status;
569
739
  }
570
740
 
571
741
  // --- Lifecycle helpers ---
@@ -582,18 +752,6 @@ export class A2AThreadRuntimeCore {
582
752
  this.setRunning(false);
583
753
  }
584
754
 
585
- private resetHead(parentId: string | null | undefined) {
586
- if (!parentId) {
587
- if (this.messages.length) {
588
- this.messages = [];
589
- }
590
- return;
591
- }
592
- const idx = this.messages.findIndex((message) => message.id === parentId);
593
- if (idx === -1) return;
594
- this.messages = this.messages.slice(0, idx + 1);
595
- }
596
-
597
755
  // --- History persistence ---
598
756
 
599
757
  private recordHistoryEntry(parentId: string | null, message: ThreadMessage) {
@@ -612,7 +770,7 @@ export class A2AThreadRuntimeCore {
612
770
  if (!this.history) return;
613
771
  const parentId = this.assistantHistoryParents.get(messageId);
614
772
  if (parentId === undefined) return;
615
- const message = this.messages.find((m) => m.id === messageId);
773
+ const message = this.tryGetMessage(messageId)?.message;
616
774
  if (!message || message.role !== "assistant") return;
617
775
  if (
618
776
  message.status?.type !== "complete" &&
@@ -245,12 +245,61 @@ describe("contentPartsToA2AParts", () => {
245
245
  expect(result).toEqual([{ text: "hi" }]);
246
246
  });
247
247
 
248
- it("converts image parts", () => {
248
+ it("converts image URL parts without stamping a mediaType", () => {
249
249
  const result = contentPartsToA2AParts([
250
250
  { type: "image", image: "https://img.com/a.png" },
251
251
  ]);
252
+ expect(result).toEqual([{ url: "https://img.com/a.png" }]);
253
+ });
254
+
255
+ it("applies the fallback MIME type to image URL parts", () => {
256
+ const result = contentPartsToA2AParts(
257
+ [{ type: "image", image: "https://img.com/a.png" }],
258
+ "image/png",
259
+ );
260
+ expect(result).toEqual([
261
+ { url: "https://img.com/a.png", mediaType: "image/png" },
262
+ ]);
263
+ });
264
+
265
+ it("applies the fallback MIME type and filename to image URL parts together", () => {
266
+ const result = contentPartsToA2AParts(
267
+ [{ type: "image", image: "https://img.com/a.png", filename: "a.png" }],
268
+ "image/png",
269
+ );
252
270
  expect(result).toEqual([
253
- { url: "https://img.com/a.png", mediaType: "image/*" },
271
+ {
272
+ url: "https://img.com/a.png",
273
+ mediaType: "image/png",
274
+ filename: "a.png",
275
+ },
276
+ ]);
277
+ });
278
+
279
+ it("passes non-base64 data URLs through as URLs for image parts", () => {
280
+ const result = contentPartsToA2AParts([
281
+ { type: "image", image: "data:text/plain,hello" },
282
+ ]);
283
+ expect(result).toEqual([{ url: "data:text/plain,hello" }]);
284
+ });
285
+
286
+ it("converts image data URLs to raw bytes with the embedded MIME type", () => {
287
+ const result = contentPartsToA2AParts([
288
+ { type: "image", image: "data:image/png;base64,aGVsbG8=" },
289
+ ]);
290
+ expect(result).toEqual([{ raw: "aGVsbG8=", mediaType: "image/png" }]);
291
+ });
292
+
293
+ it("propagates image filenames", () => {
294
+ const result = contentPartsToA2AParts([
295
+ {
296
+ type: "image",
297
+ image: "data:image/png;base64,aGVsbG8=",
298
+ filename: "a.png",
299
+ },
300
+ ]);
301
+ expect(result).toEqual([
302
+ { raw: "aGVsbG8=", mediaType: "image/png", filename: "a.png" },
254
303
  ]);
255
304
  });
256
305
 
@@ -259,6 +308,64 @@ describe("contentPartsToA2AParts", () => {
259
308
  expect(result).toEqual([]);
260
309
  });
261
310
 
311
+ it("converts file parts with http URLs", () => {
312
+ const result = contentPartsToA2AParts([
313
+ {
314
+ type: "file",
315
+ data: "https://files.com/doc.pdf",
316
+ mimeType: "application/pdf",
317
+ filename: "doc.pdf",
318
+ },
319
+ ]);
320
+ expect(result).toEqual([
321
+ {
322
+ url: "https://files.com/doc.pdf",
323
+ mediaType: "application/pdf",
324
+ filename: "doc.pdf",
325
+ },
326
+ ]);
327
+ });
328
+
329
+ it("converts file parts with data URLs to raw bytes", () => {
330
+ const result = contentPartsToA2AParts([
331
+ {
332
+ type: "file",
333
+ data: "data:application/pdf;base64,ZmlsZQ==",
334
+ mimeType: "application/pdf",
335
+ },
336
+ ]);
337
+ expect(result).toEqual([{ raw: "ZmlsZQ==", mediaType: "application/pdf" }]);
338
+ });
339
+
340
+ it("converts file parts with raw base64 data", () => {
341
+ const result = contentPartsToA2AParts([
342
+ { type: "file", data: "ZmlsZQ==", mimeType: "text/csv" },
343
+ ]);
344
+ expect(result).toEqual([{ raw: "ZmlsZQ==", mediaType: "text/csv" }]);
345
+ });
346
+
347
+ it("falls back to the attachment MIME type when the file part MIME is empty", () => {
348
+ const result = contentPartsToA2AParts(
349
+ [{ type: "file", data: "ZmlsZQ==", mimeType: "" }],
350
+ "application/pdf",
351
+ );
352
+ expect(result).toEqual([{ raw: "ZmlsZQ==", mediaType: "application/pdf" }]);
353
+ });
354
+
355
+ it("omits mediaType when no MIME type is known", () => {
356
+ const result = contentPartsToA2AParts([
357
+ { type: "file", data: "ZmlsZQ==", mimeType: "" },
358
+ ]);
359
+ expect(result).toEqual([{ raw: "ZmlsZQ==" }]);
360
+ });
361
+
362
+ it("skips file parts with no data", () => {
363
+ const result = contentPartsToA2AParts([
364
+ { type: "file", mimeType: "application/pdf" },
365
+ ]);
366
+ expect(result).toEqual([]);
367
+ });
368
+
262
369
  it("skips unknown part types", () => {
263
370
  const result = contentPartsToA2AParts([
264
371
  { type: "text", text: "hi" },
@@ -1,6 +1,7 @@
1
1
  "use client";
2
2
 
3
3
  import type { MessageStatus, ThreadAssistantMessage } from "@assistant-ui/core";
4
+ import { httpUrlPattern, parseDataUrl } from "@assistant-ui/core/internal";
4
5
  import type { A2AMessage, A2APart, A2ATaskState } from "./types";
5
6
 
6
7
  function isImageMediaType(mediaType?: string): boolean {
@@ -89,18 +90,59 @@ export function taskStateToMessageStatus(state: A2ATaskState): MessageStatus {
89
90
  export function contentPartsToA2AParts(
90
91
  content: ReadonlyArray<{
91
92
  type: string;
92
- text?: string;
93
- image?: string;
93
+ text?: string | undefined;
94
+ image?: string | undefined;
95
+ data?: string | undefined;
96
+ mimeType?: string | undefined;
97
+ filename?: string | undefined;
94
98
  }>,
99
+ fallbackMimeType?: string,
95
100
  ): A2APart[] {
96
101
  return content
97
102
  .map((part): A2APart | null => {
98
103
  switch (part.type) {
99
104
  case "text":
100
105
  return { text: part.text ?? "" };
101
- case "image":
106
+ case "image": {
102
107
  if (!part.image) return null;
103
- return { url: part.image, mediaType: "image/*" };
108
+ const parsed = parseDataUrl(part.image);
109
+ if (parsed) {
110
+ return {
111
+ raw: parsed.data,
112
+ mediaType: parsed.mimeType,
113
+ ...(part.filename && { filename: part.filename }),
114
+ };
115
+ }
116
+ return {
117
+ url: part.image,
118
+ ...(fallbackMimeType && { mediaType: fallbackMimeType }),
119
+ ...(part.filename && { filename: part.filename }),
120
+ };
121
+ }
122
+ case "file": {
123
+ if (!part.data) return null;
124
+ const declaredMimeType = part.mimeType || fallbackMimeType;
125
+ if (httpUrlPattern.test(part.data)) {
126
+ return {
127
+ url: part.data,
128
+ ...(declaredMimeType && { mediaType: declaredMimeType }),
129
+ ...(part.filename && { filename: part.filename }),
130
+ };
131
+ }
132
+ const parsed = parseDataUrl(part.data);
133
+ if (parsed) {
134
+ return {
135
+ raw: parsed.data,
136
+ mediaType: parsed.mimeType,
137
+ ...(part.filename && { filename: part.filename }),
138
+ };
139
+ }
140
+ return {
141
+ raw: part.data,
142
+ ...(declaredMimeType && { mediaType: declaredMimeType }),
143
+ ...(part.filename && { filename: part.filename }),
144
+ };
145
+ }
104
146
  default:
105
147
  return null;
106
148
  }