From 65668bd5699769ebd39bcb709f2d509e49edfe56 Mon Sep 17 00:00:00 2001 From: Brendan Morin Date: Thu, 4 Aug 2022 14:18:09 -0700 Subject: [PATCH 1/4] add logic to handle circular loops, and ignore deprecated fields --- pbspark/_proto.py | 34 +++++++++++++++++++++++++++++----- 1 file changed, 29 insertions(+), 5 deletions(-) diff --git a/pbspark/_proto.py b/pbspark/_proto.py index bfaf352..427955b 100644 --- a/pbspark/_proto.py +++ b/pbspark/_proto.py @@ -1,6 +1,8 @@ import inspect +import logging import typing as t from contextlib import contextmanager +from copy import copy from functools import wraps from google.protobuf import json_format @@ -237,34 +239,56 @@ def parse_dict( return parser.ConvertMessage(value=value, message=message, path=None) def get_spark_schema( - self, - descriptor: t.Union[t.Type[Message], Descriptor], - options: t.Optional[dict] = None, + self, + descriptor: t.Union[t.Type[Message], Descriptor], + options: t.Optional[dict] = None, + _seen_descriptors: t.Optional[set] = None ) -> DataType: """Generate a spark schema from a message type or descriptor - Given a message type generated from protoc (or its descriptor), create a spark schema derived from the protobuf schema when serializing with ``MessageToDict``. """ + # track which descriptors have been seen in the current proto graph (for loop identification) + _seen_descriptors_ = copy(_seen_descriptors) or set() + options = options or {} use_camelcase = not options.get("preserving_proto_field_name", False) + ignore_deprecated = options.get("ignore_deprecated", False) + schema = [] if inspect.isclass(descriptor) and issubclass(descriptor, Message): descriptor_ = descriptor.DESCRIPTOR else: descriptor_ = descriptor # type: ignore[assignment] + _seen_descriptors_.add(descriptor_.full_name) + full_name = descriptor_.full_name if full_name in self._message_type_to_spark_type_map: return self._message_type_to_spark_type_map[full_name] + for field in descriptor_.fields: + field_full_name = field.message_type.full_name + if field.has_options: + field_options = field.GetOptions() + + # Optionally ignore deprecated fields + if field_options.deprecated and ignore_deprecated: + continue + + # Check for recursive loops in proto definition + if field.message_type != None: # noqa ("is None" is not the same as "!= None" here) + if field_full_name in _seen_descriptors_: + logging.warning(f"Circular protobuf definition detected! Ignoring field: {field_full_name}") + continue + spark_type: DataType if field.cpp_type == FieldDescriptor.CPPTYPE_MESSAGE: full_name = field.message_type.full_name if full_name in self._message_type_to_spark_type_map: spark_type = self._message_type_to_spark_type_map[full_name] else: - spark_type = self.get_spark_schema(field.message_type, options) + spark_type = self.get_spark_schema(field.message_type, options, _seen_descriptors_) # protobuf converts to/from b64 strings, but we prefer to stay as bytes elif ( field.cpp_type == FieldDescriptor.CPPTYPE_STRING From 9766e9750800bc06eb54fac87d26a835a1935ecb Mon Sep 17 00:00:00 2001 From: Brendan Morin Date: Thu, 4 Aug 2022 14:19:06 -0700 Subject: [PATCH 2/4] fix indentation --- pbspark/_proto.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/pbspark/_proto.py b/pbspark/_proto.py index 427955b..33d194d 100644 --- a/pbspark/_proto.py +++ b/pbspark/_proto.py @@ -239,10 +239,10 @@ def parse_dict( return parser.ConvertMessage(value=value, message=message, path=None) def get_spark_schema( - self, - descriptor: t.Union[t.Type[Message], Descriptor], - options: t.Optional[dict] = None, - _seen_descriptors: t.Optional[set] = None + self, + descriptor: t.Union[t.Type[Message], Descriptor], + options: t.Optional[dict] = None, + _seen_descriptors: t.Optional[set] = None ) -> DataType: """Generate a spark schema from a message type or descriptor Given a message type generated from protoc (or its descriptor), From 00a376a84e3fe97bc9bf166dc47dc9929b8b92e2 Mon Sep 17 00:00:00 2001 From: Brendan Morin Date: Thu, 4 Aug 2022 14:23:03 -0700 Subject: [PATCH 3/4] add option to fail on loop --- pbspark/_proto.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/pbspark/_proto.py b/pbspark/_proto.py index 33d194d..2e3d2d2 100644 --- a/pbspark/_proto.py +++ b/pbspark/_proto.py @@ -255,6 +255,7 @@ def get_spark_schema( options = options or {} use_camelcase = not options.get("preserving_proto_field_name", False) ignore_deprecated = options.get("ignore_deprecated", False) + fail_on_loop = options.get("fail_on_loop", False) schema = [] if inspect.isclass(descriptor) and issubclass(descriptor, Message): @@ -279,8 +280,11 @@ def get_spark_schema( # Check for recursive loops in proto definition if field.message_type != None: # noqa ("is None" is not the same as "!= None" here) if field_full_name in _seen_descriptors_: - logging.warning(f"Circular protobuf definition detected! Ignoring field: {field_full_name}") - continue + if fail_on_loop: + logging.warning(f"Circular protobuf definition detected! Ignoring field: {field_full_name}") + continue + else: + raise ValueError(f"Circular protobuf definition detected: {field_full_name}") spark_type: DataType if field.cpp_type == FieldDescriptor.CPPTYPE_MESSAGE: From 171ff45bfab373f4011ba214c7b5f57927df8297 Mon Sep 17 00:00:00 2001 From: Brendan Morin Date: Thu, 4 Aug 2022 14:24:06 -0700 Subject: [PATCH 4/4] minor modification to circular definition naming --- pbspark/_proto.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pbspark/_proto.py b/pbspark/_proto.py index 2e3d2d2..7cc5fdc 100644 --- a/pbspark/_proto.py +++ b/pbspark/_proto.py @@ -255,7 +255,7 @@ def get_spark_schema( options = options or {} use_camelcase = not options.get("preserving_proto_field_name", False) ignore_deprecated = options.get("ignore_deprecated", False) - fail_on_loop = options.get("fail_on_loop", False) + ignore_circular_definitions = options.get("ignore_circular_definitions", False) schema = [] if inspect.isclass(descriptor) and issubclass(descriptor, Message): @@ -280,7 +280,7 @@ def get_spark_schema( # Check for recursive loops in proto definition if field.message_type != None: # noqa ("is None" is not the same as "!= None" here) if field_full_name in _seen_descriptors_: - if fail_on_loop: + if ignore_circular_definitions: logging.warning(f"Circular protobuf definition detected! Ignoring field: {field_full_name}") continue else: