Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
188 changes: 121 additions & 67 deletions pybess/protobuf_to_dict.py
Original file line number Diff line number Diff line change
Expand Up @@ -153,84 +153,138 @@ def dict_to_protobuf(pb_klass_or_instance, values, type_callable_map=REVERSE_TYP
return _dict_to_protobuf(instance, values, type_callable_map, strict)


def _get_field_mapping(pb, dict_value, strict):
field_mapping = []
for key, value in dict_value.items():
if key == EXTENSION_CONTAINER:
continue
if key not in pb.DESCRIPTOR.fields_by_name:
if strict:
raise KeyError("%s does not have a field called %s" %
(pb.__class__.__name__, key))
continue
field_mapping.append((pb.DESCRIPTOR.fields_by_name[
def _process_regular_fields(pb, dict_value, strict, field_mapping):
"""Process regular (non-extension) fields from the dictionary."""
for key, value in dict_value.items():
if key == EXTENSION_CONTAINER:
continue
if key not in pb.DESCRIPTOR.fields_by_name:
if strict:
raise KeyError("%s does not have a field called %s" %
(pb.__class__.__name__, key))
continue
field_mapping.append((pb.DESCRIPTOR.fields_by_name[
key], value, getattr(pb, key, None)))

for ext_num, ext_val in dict_value.get(EXTENSION_CONTAINER, {}).items():
try:
ext_num = int(ext_num)
except ValueError:
raise ValueError("Extension keys must be integers.")
if ext_num not in pb._extensions_by_number:
if strict:
raise KeyError("%s does not have a extension with number %s. "
"Perhaps you forgot to import it?" %
(pb.__class__.__name__, key))
continue
ext_field = pb._extensions_by_number[ext_num]
pb_val = pb.Extensions[ext_field]

def _process_extension_fields(pb, dict_value, strict, field_mapping):
"""Process extension fields from the dictionary."""
for ext_num, ext_val in dict_value.get(EXTENSION_CONTAINER, {}).items():
try:
ext_num = int(ext_num)
except ValueError:
raise ValueError("Extension keys must be integers.")

if ext_num not in pb._extensions_by_number:
if strict:
raise KeyError("%s does not have a extension with number %s. "
"Perhaps you forgot to import it?" %
(pb.__class__.__name__, ext_num))
continue

ext_field = pb._extensions_by_number[ext_num]
pb_val = pb.Extensions[ext_field]
field_mapping.append((ext_field, ext_val, pb_val))

def _get_field_mapping(pb, dict_value, strict):
"""Get field mapping for dictionary to protobuf conversion."""
field_mapping = []

# Process regular fields
_process_regular_fields(pb, dict_value, strict, field_mapping)

# Process extension fields
_process_extension_fields(pb, dict_value, strict, field_mapping)

return field_mapping


def _dict_to_protobuf(pb, value, type_callable_map, strict):
fields = _get_field_mapping(pb, value, strict)

if sys.version_info[0] == 2:
basestr = basestring
else:
basestr = str

for field, input_value, pb_value in fields:
if field.is_repeated:
if field.message_type and field.message_type.has_options and \
field.message_type.GetOptions().map_entry:
# Special processing for nested dict
if isinstance(input_value, dict) and all([isinstance(x, dict) for x in input_value.values()]):
for k, v in input_value.items():
_dict_to_protobuf(
pb_value[k], input_value[k], type_callable_map, strict)
else:
pb_value.update(input_value)
continue
for item in input_value:
if field.type == FieldDescriptor.TYPE_MESSAGE:
m = pb_value.add()
_dict_to_protobuf(m, item, type_callable_map, strict)
elif field.type == FieldDescriptor.TYPE_ENUM and isinstance(item, basestr):
pb_value.append(_string_to_enum(field, item))
else:
pb_value.append(item)
continue
if field.type == FieldDescriptor.TYPE_MESSAGE:
_dict_to_protobuf(pb_value, input_value, type_callable_map, strict)
continue

def _dict_to_protobuf(pb, value, type_callable_map, strict):
"""Populates a protobuf message from a dictionary."""
fields = _get_field_mapping(pb, value, strict)
basestr = _get_base_string_type()

for field, input_value, pb_value in fields:
# 1. Handle repeated fields (exits loop on match)
if _handle_repeated_field(field, input_value, pb_value, type_callable_map, strict, basestr):
continue

# 2. Handle nested message fields (exits loop on match)
if _handle_message_field(field, input_value, pb_value, type_callable_map, strict):
continue

# 3. Apply type callable mapping (must run before extension checks)
if field.type in type_callable_map:
input_value = type_callable_map[field.type](input_value)

if field.is_extension:
pb.Extensions[field] = input_value
continue


# 4. Handle extensions (exits loop on match)
if _handle_extension_field(field, input_value, pb):
continue

# 5. Apply enum conversion (only runs on standard fields)
if field.type == FieldDescriptor.TYPE_ENUM and isinstance(input_value, basestr):
input_value = _string_to_enum(field, input_value)

setattr(pb, field.name, input_value)


# 6. Set standard attribute
setattr(pb, field.name, input_value)

return pb

def _get_base_string_type():
"""Get the appropriate string type for the Python version."""
return basestring if sys.version_info[0] == 2 else str

def _handle_repeated_field(field, input_value, pb_value, type_callable_map, strict, basestr):
"""Handle repeated fields including map entries."""
if field.label != FieldDescriptor.LABEL_REPEATED:
return False

if _is_map_field(field):
_handle_map_field(field, input_value, pb_value, type_callable_map, strict)
else:
# Added 'strict' to this call
_handle_regular_repeated_field(field, input_value, pb_value, type_callable_map, strict, basestr)

return True

def _is_map_field(field):
"""Check if field is a map entry."""
return (field.message_type and
field.message_type.has_options and
field.message_type.GetOptions().map_entry)

def _handle_map_field(field, input_value, pb_value, type_callable_map, strict):
"""Handle map field processing."""
if isinstance(input_value, dict) and all(isinstance(x, dict) for x in input_value.values()):
for k, v in input_value.items():
_dict_to_protobuf(pb_value[k], v, type_callable_map, strict)
else:
pb_value.update(input_value)

def _handle_regular_repeated_field(field, input_value, pb_value, type_callable_map, strict, basestr):
"""Handle regular repeated field processing."""
for item in input_value:
if field.type == FieldDescriptor.TYPE_MESSAGE:
m = pb_value.add()
_dict_to_protobuf(m, item, type_callable_map, strict)
elif field.type == FieldDescriptor.TYPE_ENUM and isinstance(item, basestr):
pb_value.append(_string_to_enum(field, item))
else:
pb_value.append(item)

def _handle_message_field(field, input_value, pb_value, type_callable_map, strict):
"""Handle message field processing."""
if field.type != FieldDescriptor.TYPE_MESSAGE:
return False

_dict_to_protobuf(pb_value, input_value, type_callable_map, strict)
return True

def _handle_extension_field(field, input_value, pb):
"""Handle extension field processing."""
if not field.is_extension:
return False

pb.Extensions[field] = input_value
return True

def _string_to_enum(field, input_value):
enum_dict = field.enum_type.values_by_name
Expand Down