diff --git a/src/a2a/utils/signing.py b/src/a2a/utils/signing.py index c85a80072..96566167e 100644 --- a/src/a2a/utils/signing.py +++ b/src/a2a/utils/signing.py @@ -3,6 +3,7 @@ from collections.abc import Callable from typing import Any, TypedDict +from google.protobuf.descriptor import Descriptor, FieldDescriptor from google.protobuf.json_format import MessageToDict @@ -22,6 +23,7 @@ from a2a.types import AgentCard, AgentCardSignature from a2a.utils._jcs import MAX_DEPTH, CanonicalizationError, canonicalize +from a2a.utils.proto_utils import _field_is_repeated class SignatureVerificationError(Exception): @@ -195,6 +197,72 @@ def _clean_empty(d: Any, depth: int = 0) -> Any: return d +def _is_map(field: FieldDescriptor) -> bool: + """Returns True if the field is a protobuf map.""" + message_type = field.message_type + return message_type is not None and message_type.GetOptions().map_entry + + +def _is_well_known(descriptor: Descriptor | Any) -> bool: + """Returns True for `google.protobuf` types, which carry free-form JSON.""" + return descriptor.full_name.startswith('google.protobuf.') + + +def _clean_field(value: Any, field: FieldDescriptor, depth: int) -> Any: + """Removes empty values from the JSON form of one message field.""" + message_type = field.message_type + if message_type is None or _is_well_known(message_type): + return _clean_empty(value, depth) + if _is_map(field): + value_type = message_type.fields_by_name['value'].message_type + if value_type is None or _is_well_known(value_type): + return _clean_empty(value, depth) + cleaned_map = { + k: cleaned_v + for k, v in value.items() + if (cleaned_v := _clean_message(v, value_type, depth + 1)) + } + return cleaned_map or None + if _field_is_repeated(field): + cleaned_list = [ + cleaned_v + for v in value + if (cleaned_v := _clean_message(v, message_type, depth + 1)) + ] + return cleaned_list or None + return _clean_message(value, message_type, depth) or None + + +def _clean_message( + message_dict: dict[str, Any], + descriptor: Descriptor | Any, + depth: int = 0, +) -> dict[str, Any]: + """Removes empty values from the JSON form of a message, by descriptor. + + `message_dict` is the `MessageToDict` output for a message of type + `descriptor`. Walking the descriptor alongside the JSON keeps the field + each value belongs to known at every level, which `_clean_empty` alone + cannot tell. Free-form values (`google.protobuf.Struct` and friends) and + keys the descriptor does not know fall back to `_clean_empty`. + """ + if depth > MAX_DEPTH: + raise CanonicalizationError( + f'nesting exceeds the maximum depth of {MAX_DEPTH}' + ) + fields = {field.json_name: field for field in descriptor.fields} + cleaned: dict[str, Any] = {} + for key, value in message_dict.items(): + field = fields.get(key) + if field is None: + cleaned_value = _clean_empty(value, depth + 1) + else: + cleaned_value = _clean_field(value, field, depth + 1) + if cleaned_value is not None: + cleaned[key] = cleaned_value + return cleaned + + def _canonicalize_agent_card(agent_card: AgentCard) -> str: """Canonicalizes the Agent Card JSON according to RFC 8785 (JCS).""" card_dict = MessageToDict( @@ -203,6 +271,6 @@ def _canonicalize_agent_card(agent_card: AgentCard) -> str: # Remove signatures field if present card_dict.pop('signatures', None) - # Recursively remove empty values - cleaned_dict = _clean_empty(card_dict) - return canonicalize(cleaned_dict) + # Remove empty values, walking the AgentCard descriptor + cleaned_dict = _clean_message(card_dict, AgentCard.DESCRIPTOR) + return canonicalize(cleaned_dict or None) diff --git a/tests/utils/test_signing.py b/tests/utils/test_signing.py index 616aab9f3..fd8939815 100644 --- a/tests/utils/test_signing.py +++ b/tests/utils/test_signing.py @@ -3,14 +3,24 @@ import pytest from a2a.types.a2a_pb2 import ( + APIKeySecurityScheme, AgentCapabilities, AgentCard, AgentCardSignature, + AgentExtension, AgentInterface, + AgentProvider, AgentSkill, + AuthorizationCodeOAuthFlow, + OAuth2SecurityScheme, + OAuthFlows, + SecurityRequirement, + SecurityScheme, + StringList, ) from a2a.utils import signing from cryptography.hazmat.primitives.asymmetric import ec +from google.protobuf.json_format import MessageToDict from jwt.utils import base64url_encode @@ -287,3 +297,110 @@ def test_clean_empty_does_not_mutate_input(): signing._clean_empty(original) assert original == original_copy + + +@pytest.fixture +def full_agent_card() -> AgentCard: + """A card that exercises nested messages, maps and free-form values.""" + card = AgentCard( + name='Full Agent', + description='A card that exercises nested messages', + supported_interfaces=[ + AgentInterface( + url='https://example.com/a2a/v1', + protocol_binding='JSONRPC', + protocol_version='1.0', + tenant='', + ) + ], + provider=AgentProvider( + url='https://example.com', organization='Example' + ), + version='1.0.0', + capabilities=AgentCapabilities( + streaming=False, + extensions=[ + AgentExtension(uri='https://example.com/ext/1'), + AgentExtension(uri='https://example.com/ext/2', description=''), + ], + ), + security_schemes={ + 'key': SecurityScheme( + api_key_security_scheme=APIKeySecurityScheme( + location='header', name='X-API-Key', description='' + ) + ), + 'oauth': SecurityScheme( + oauth2_security_scheme=OAuth2SecurityScheme( + flows=OAuthFlows( + authorization_code=AuthorizationCodeOAuthFlow( + authorization_url='https://example.com/auth', + token_url='https://example.com/token', + scopes={'read': 'Read access'}, + ) + ) + ) + ), + }, + security_requirements=[ + SecurityRequirement(schemes={'oauth': StringList(list=['read'])}), + SecurityRequirement(), + ], + default_input_modes=['text/plain'], + default_output_modes=['text/plain'], + skills=[ + AgentSkill( + id='skill1', + name='Skill', + description='A skill', + tags=['test'], + examples=[], + ) + ], + icon_url='', + ) + card.capabilities.extensions[0].params.update( + {'empty': '', 'nested': {'list': [], 'kept': 0}} + ) + return card + + +def test_clean_message_matches_clean_empty_on_full_card( + full_agent_card: AgentCard, +): + """Descriptor-aware cleaning gives the same result as `_clean_empty`.""" + card_dict = MessageToDict(full_agent_card) + assert signing._clean_message( + card_dict, AgentCard.DESCRIPTOR + ) == signing._clean_empty(card_dict) + + +def test_canonicalize_full_card_prunes_optional_defaults( + full_agent_card: AgentCard, +): + """Optional and free-form empty values are pruned at every level.""" + result = signing._canonicalize_agent_card(full_agent_card) + assert '"tenant"' not in result + assert '"iconUrl"' not in result + assert '"examples"' not in result + assert '"empty"' not in result + assert '"list"' in result # the StringList inside the security requirement + assert '"params":{"nested":{"kept":0}}' in result + assert '"streaming":false' in result + # The empty SecurityRequirement element is dropped, the other one stays. + assert ( + '"securityRequirements":[{"schemes":{"oauth":{"list":["read"]}}}]' + in result + ) + + +def test_clean_message_bounds_depth(): + """Descriptor-aware cleaning keeps the depth bound of `_clean_empty`.""" + nested: dict[str, Any] = {} + cursor = nested + for _ in range(signing.MAX_DEPTH + 5): + cursor['a'] = {} + cursor = cursor['a'] + card_dict = {'capabilities': {'extensions': [{'params': nested}]}} + with pytest.raises(signing.CanonicalizationError): + signing._clean_message(card_dict, AgentCard.DESCRIPTOR)