Skip to content

Commit a2db488

Browse files
committed
bugfix: 修复ag-ui协议返回message_id为空问题
Made-with: Cursor
1 parent 043fcfb commit a2db488

2 files changed

Lines changed: 86 additions & 41 deletions

File tree

tests/server/a2a/converters/test_event_converter.py

Lines changed: 14 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -301,9 +301,10 @@ def test_includes_optional_fields(self):
301301
class TestBuildMessageMetadata:
302302
def test_includes_object_type_and_tag(self):
303303
event = _make_event(text="hi", object_type="chat.completion", tag="my_tag")
304-
meta = _build_message_metadata(event)
304+
meta = _build_message_metadata(event, "eff-1")
305305
assert meta[MESSAGE_METADATA_OBJECT_TYPE_KEY] == "chat.completion"
306306
assert meta[MESSAGE_METADATA_TAG_KEY] == "my_tag"
307+
assert meta[MESSAGE_METADATA_RESPONSE_ID_KEY] == "eff-1"
307308

308309

309310
# ---------------------------------------------------------------------------
@@ -337,12 +338,12 @@ def test_does_nothing_without_long_running_ids(self):
337338
class TestBuildMessage:
338339
def test_returns_none_for_empty_parts(self):
339340
event = _make_event(text="hi")
340-
assert _build_message(event, [], Role.agent) is None
341+
assert _build_message(event, [], Role.agent, "e1") is None
341342

342343
def test_returns_message_with_parts(self):
343344
event = _make_event(text="hi", response_id="resp-1")
344345
parts = [A2APart(root=TextPart(text="hi"))]
345-
msg = _build_message(event, parts, Role.agent)
346+
msg = _build_message(event, parts, Role.agent, "resp-1")
346347
assert msg is not None
347348
assert msg.role == Role.agent
348349
assert msg.message_id == "resp-1"
@@ -636,7 +637,7 @@ def test_basic_working(self):
636637
)
637638
event = _make_event(text="hi")
638639
ctx = _make_invocation_context()
639-
result = _create_status_update_event(msg, ctx, event, "t1", "ctx1")
640+
result = _create_status_update_event(msg, ctx, event, "t1", "ctx1", effective_id="m1")
640641
assert result.status.state == TaskState.working
641642

642643
def test_auth_required_for_euc(self):
@@ -650,7 +651,7 @@ def test_auth_required_for_euc(self):
650651
msg = Message(message_id="m1", role=Role.agent, parts=[A2APart(root=dp)])
651652
event = _make_event(function_call=FunctionCall(name=REQUEST_EUC_FUNCTION_CALL_NAME, args={}))
652653
ctx = _make_invocation_context()
653-
result = _create_status_update_event(msg, ctx, event, "t1", "ctx1")
654+
result = _create_status_update_event(msg, ctx, event, "t1", "ctx1", effective_id="m1")
654655
assert result.status.state == TaskState.auth_required
655656

656657
def test_input_required_for_long_running(self):
@@ -664,7 +665,7 @@ def test_input_required_for_long_running(self):
664665
msg = Message(message_id="m1", role=Role.agent, parts=[A2APart(root=dp)])
665666
event = _make_event(function_call=FunctionCall(name="other_tool", args={}))
666667
ctx = _make_invocation_context()
667-
result = _create_status_update_event(msg, ctx, event, "t1", "ctx1")
668+
result = _create_status_update_event(msg, ctx, event, "t1", "ctx1", effective_id="m1")
668669
assert result.status.state == TaskState.input_required
669670

670671

@@ -680,15 +681,19 @@ def test_basic(self):
680681
)
681682
event = _make_event(text="hi", response_id="resp-1")
682683
ctx = _make_invocation_context()
683-
result = _create_artifact_update_event(msg, event, ctx, task_id="t1", context_id="ctx1")
684-
assert result.artifact.artifact_id == "resp-1"
684+
result = _create_artifact_update_event(
685+
msg, event, ctx, task_id="t1", context_id="ctx1", effective_id="m1"
686+
)
687+
assert result.artifact.artifact_id == "m1"
685688
assert result.last_chunk is False
686689

687690
def test_last_chunk(self):
688691
msg = Message(message_id="m1", role=Role.agent, parts=[A2APart(root=TextPart(text="hi"))])
689692
event = _make_event(text="hi")
690693
ctx = _make_invocation_context()
691-
result = _create_artifact_update_event(msg, event, ctx, task_id="t1", context_id="ctx1", last_chunk=True)
694+
result = _create_artifact_update_event(
695+
msg, event, ctx, task_id="t1", context_id="ctx1", last_chunk=True, effective_id="m1"
696+
)
692697
assert result.last_chunk is True
693698
assert result.artifact.artifact_id == ""
694699
assert result.artifact.parts == []

trpc_agent_sdk/server/a2a/converters/_event_converter.py

Lines changed: 72 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -206,24 +206,25 @@ def _build_context_metadata(event: Event, ctx: InvocationContext) -> Dict[str, A
206206
return metadata
207207

208208

209-
def _build_message_metadata(event: Event) -> Dict[str, Any]:
209+
def _build_message_metadata(event: Event, effective_id: str) -> Dict[str, Any]:
210210
"""Build message/event metadata (object_type, tag, llm_response_id)."""
211211
return {
212212
MESSAGE_METADATA_OBJECT_TYPE_KEY: _infer_message_object_type(event) or "",
213213
MESSAGE_METADATA_TAG_KEY: _infer_message_tag(event),
214-
MESSAGE_METADATA_RESPONSE_ID_KEY: event.response_id or "",
214+
MESSAGE_METADATA_RESPONSE_ID_KEY: effective_id,
215215
}
216216

217217

218-
def _build_event_metadata(event: Event, message: Message, ctx: InvocationContext) -> Dict[str, Any]:
218+
def _build_event_metadata(event: Event, message: Message, ctx: InvocationContext, effective_id: str) -> Dict[str, Any]:
219219
metadata = _build_context_metadata(event, ctx)
220-
msg_meta = _build_message_metadata(event)
220+
msg_meta = _build_message_metadata(event, effective_id)
221221
set_metadata(metadata, MESSAGE_METADATA_OBJECT_TYPE_KEY, msg_meta.get(MESSAGE_METADATA_OBJECT_TYPE_KEY) or "")
222222
set_metadata(metadata, MESSAGE_METADATA_TAG_KEY, msg_meta.get(MESSAGE_METADATA_TAG_KEY) or "")
223223
set_metadata(metadata, MESSAGE_METADATA_RESPONSE_ID_KEY, msg_meta.get(MESSAGE_METADATA_RESPONSE_ID_KEY) or "")
224+
streaming_delta = A2A_DATA_PART_METADATA_TYPE_STREAMING_FUNCTION_CALL_DELTA
224225
if any(
225-
get_metadata(p.root.metadata, A2A_DATA_PART_METADATA_TYPE_KEY) ==
226-
A2A_DATA_PART_METADATA_TYPE_STREAMING_FUNCTION_CALL_DELTA for p in message.parts if p.root.metadata):
226+
get_metadata(p.root.metadata, A2A_DATA_PART_METADATA_TYPE_KEY) == streaming_delta for p in message.parts
227+
if p.root.metadata):
227228
set_metadata(metadata, "streaming_tool_call", "true")
228229
return metadata
229230

@@ -234,27 +235,41 @@ def _mark_long_running_tools(a2a_parts: List[A2APart], event: Event) -> None:
234235
return
235236
for a2a_part in a2a_parts:
236237
root = a2a_part.root
237-
if (isinstance(root, DataPart) and root.metadata and get_metadata(
238-
root.metadata, A2A_DATA_PART_METADATA_TYPE_KEY) == A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL
239-
and root.data.get("id") in event.long_running_tool_ids):
240-
set_metadata(root.metadata, A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY, True)
238+
if not isinstance(root, DataPart) or not root.metadata:
239+
continue
240+
if get_metadata(root.metadata, A2A_DATA_PART_METADATA_TYPE_KEY) != A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL:
241+
continue
242+
if root.data.get("id") not in event.long_running_tool_ids:
243+
continue
244+
set_metadata(root.metadata, A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY, True)
245+
246+
247+
def _effective_response_id(event: Event) -> str:
248+
"""Return ``response_id`` when present, otherwise a new UUID.
249+
250+
Callers that need the same id across multiple locations should invoke this
251+
once and pass the result explicitly.
252+
"""
253+
return event.response_id or str(uuid.uuid4())
241254

242255

243-
def _build_message(event: Event, a2a_parts: List[A2APart], role: Role) -> Optional[Message]:
256+
def _build_message(event: Event, a2a_parts: List[A2APart], role: Role, effective_id: str) -> Optional[Message]:
244257
"""Assemble an A2A Message from converted parts, or return None if empty."""
245258
if not a2a_parts:
246259
return None
247-
message_id = event.response_id or str(uuid.uuid4())
248-
message = Message(message_id=message_id, role=role, parts=a2a_parts)
249-
msg_meta = _build_message_metadata(event)
260+
message = Message(message_id=effective_id, role=role, parts=a2a_parts)
261+
msg_meta = _build_message_metadata(event, effective_id)
250262
if msg_meta:
251263
message.metadata = msg_meta
252264
return message
253265

254266

255267
def _is_streaming_delta(a2a_part: A2APart) -> bool:
256-
return (a2a_part.root.metadata is not None and get_metadata(a2a_part.root.metadata, A2A_DATA_PART_METADATA_TYPE_KEY)
257-
== A2A_DATA_PART_METADATA_TYPE_STREAMING_FUNCTION_CALL_DELTA)
268+
meta = a2a_part.root.metadata
269+
if meta is None:
270+
return False
271+
t = get_metadata(meta, A2A_DATA_PART_METADATA_TYPE_KEY)
272+
return t == A2A_DATA_PART_METADATA_TYPE_STREAMING_FUNCTION_CALL_DELTA
258273

259274

260275
def _collect_parts(
@@ -320,7 +335,8 @@ def convert_event_to_a2a_message(
320335
return None
321336

322337
a2a_parts = _collect_parts(event, **rules)
323-
return _build_message(event, a2a_parts, role)
338+
effective_id = _effective_response_id(event)
339+
return _build_message(event, a2a_parts, role, effective_id)
324340

325341

326342
def convert_content_to_a2a_message(
@@ -425,8 +441,8 @@ def convert_a2a_message_to_event(
425441
if gpart is None:
426442
logger.warning("Failed to convert A2A part, skipping: %s", a2a_part)
427443
continue
428-
if (metadata_is_true(a2a_part.root.metadata, A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY)
429-
and gpart.function_call):
444+
is_lr = metadata_is_true(a2a_part.root.metadata, A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY)
445+
if is_lr and gpart.function_call:
430446
long_running_tool_ids.add(gpart.function_call.id)
431447
parts.append(gpart)
432448
except Exception as ex: # pylint: disable=broad-except
@@ -436,8 +452,8 @@ def convert_a2a_message_to_event(
436452
if not parts:
437453
logger.warning("No parts could be converted from A2A message %s", a2a_message)
438454

439-
object_type = (get_metadata(msg_meta, MESSAGE_METADATA_OBJECT_TYPE_KEY)
440-
or _infer_a2a_message_object_type(parts, partial=partial) or _default_object_type(partial))
455+
ot = get_metadata(msg_meta, MESSAGE_METADATA_OBJECT_TYPE_KEY)
456+
object_type = ot or _infer_a2a_message_object_type(parts, partial=partial) or _default_object_type(partial)
441457

442458
return Event(
443459
invocation_id=inv_id,
@@ -590,31 +606,51 @@ def _create_error_status_event(
590606
)
591607

592608

609+
def _a2a_part_requests_euc_auth(part: A2APart) -> bool:
610+
root = part.root
611+
md = root.metadata
612+
if not md:
613+
return False
614+
t = get_metadata(md, A2A_DATA_PART_METADATA_TYPE_KEY)
615+
return all([
616+
t == A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL,
617+
metadata_is_true(md, A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY),
618+
root.data.get("name") == REQUEST_EUC_FUNCTION_CALL_NAME,
619+
])
620+
621+
622+
def _a2a_part_is_long_running_function_call(part: A2APart) -> bool:
623+
root = part.root
624+
md = root.metadata
625+
if not md:
626+
return False
627+
t = get_metadata(md, A2A_DATA_PART_METADATA_TYPE_KEY)
628+
return all([
629+
t == A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL,
630+
metadata_is_true(md, A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY),
631+
])
632+
633+
593634
def _create_status_update_event(
594635
message: Message,
595636
ctx: InvocationContext,
596637
event: Event,
597638
task_id: Optional[str],
598639
context_id: Optional[str],
640+
effective_id: str = "",
599641
) -> TaskStatusUpdateEvent:
600642
status = TaskStatus(state=TaskState.working, message=message, timestamp=_now_iso())
601643

602-
if any(
603-
get_metadata(p.root.metadata, A2A_DATA_PART_METADATA_TYPE_KEY) == A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL
604-
and metadata_is_true(p.root.metadata, A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY)
605-
and p.root.data.get("name") == REQUEST_EUC_FUNCTION_CALL_NAME for p in message.parts if p.root.metadata):
644+
if any(_a2a_part_requests_euc_auth(p) for p in message.parts):
606645
status.state = TaskState.auth_required
607-
elif any(
608-
get_metadata(p.root.metadata, A2A_DATA_PART_METADATA_TYPE_KEY) == A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL
609-
and metadata_is_true(p.root.metadata, A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY) for p in message.parts
610-
if p.root.metadata):
646+
elif any(_a2a_part_is_long_running_function_call(p) for p in message.parts):
611647
status.state = TaskState.input_required
612648

613649
return TaskStatusUpdateEvent(
614650
task_id=task_id,
615651
context_id=context_id,
616652
status=status,
617-
metadata=_build_event_metadata(event, message, ctx),
653+
metadata=_build_event_metadata(event, message, ctx, effective_id),
618654
final=False,
619655
)
620656

@@ -626,8 +662,9 @@ def _create_artifact_update_event(
626662
task_id: Optional[str] = None,
627663
context_id: Optional[str] = None,
628664
last_chunk: bool = False,
665+
effective_id: str = "",
629666
) -> TaskArtifactUpdateEvent:
630-
artifact_id = "" if last_chunk else (event.response_id or "")
667+
artifact_id = "" if last_chunk else effective_id
631668
return TaskArtifactUpdateEvent(
632669
task_id=task_id,
633670
context_id=context_id,
@@ -636,7 +673,7 @@ def _create_artifact_update_event(
636673
parts=[] if last_chunk else message.parts,
637674
),
638675
last_chunk=last_chunk,
639-
metadata=_build_event_metadata(event, message, ctx),
676+
metadata=_build_event_metadata(event, message, ctx, effective_id),
640677
)
641678

642679

@@ -674,12 +711,14 @@ def _notify(evt: A2AEvent) -> None:
674711

675712
message = convert_event_to_a2a_message(event, invocation_context)
676713
if message:
714+
effective_id = message.message_id
677715
status_event = _create_status_update_event(
678716
message,
679717
invocation_context,
680718
event,
681719
task_id,
682720
context_id,
721+
effective_id=effective_id,
683722
)
684723
_notify(status_event)
685724

@@ -691,6 +730,7 @@ def _notify(evt: A2AEvent) -> None:
691730
task_id=task_id,
692731
context_id=context_id,
693732
last_chunk=False,
733+
effective_id=effective_id,
694734
)
695735
a2a_events.append(artifact_event)
696736

0 commit comments

Comments
 (0)