diff --git a/dpsynth/adapters/protobuf.py b/dpsynth/adapters/protobuf.py index 85fae73a..06389d5f 100644 --- a/dpsynth/adapters/protobuf.py +++ b/dpsynth/adapters/protobuf.py @@ -16,8 +16,11 @@ from __future__ import annotations from collections.abc import Mapping, Sequence +import functools +import operator from typing import Any, Literal +from absl import logging from dpsynth import domain from google.protobuf import descriptor from google.protobuf import message @@ -107,10 +110,12 @@ def infer_domain( for field in msg_desc.fields: if _is_repeated(field): if ignore_unsupported_fields: + logging.info("Skipping repeated field %s", field.name) continue raise ValueError(f"Repeated field '{field.name}' is not supported.") if field.type == descriptor.FieldDescriptor.TYPE_MESSAGE: + # Strip the submessage prefix from numerical_bounds before recursing. prefix = f"{field.name}." sub_bounds = { k.removeprefix(prefix): v @@ -138,6 +143,7 @@ def infer_domain( dtype = "int" if field.type in _INTEGER_TYPES else "float" if field.name not in numerical_bounds: if ignore_unsupported_fields: + logging.info("Skipping numeric field %s", field.name) continue raise ValueError( f"Numerical bounds must be specified for field '{field.name}'." @@ -154,6 +160,7 @@ def infer_domain( ) else: if ignore_unsupported_fields: + logging.info("Skipping unsupported field %s", field.name) continue raise ValueError( f"Field '{field.name}' of type {field.type} is not supported." @@ -179,6 +186,7 @@ def infer_schema( return domain.Schema(attributes) +@functools.lru_cache(maxsize=64) def _leaf_field_paths(desc): """Returns dot-separated field paths for all non-repeated leaf fields.""" paths = [] @@ -190,7 +198,7 @@ def _leaf_field_paths(desc): paths.extend(f"{field.name}.{p}" for p in sub_paths) else: paths.append(field.name) - return paths + return tuple(paths) def _resolve_leaf(msg, path): @@ -201,6 +209,30 @@ def _resolve_leaf(msg, path): return msg, leaf +@functools.lru_cache(maxsize=64) +def _compile_getters(desc, paths, enum_format): + """Precompiles per-field getters for fast tuple extraction.""" + getters = [] + for path in paths: + # Walk nested message descriptors to inspect the leaf FieldDescriptor. + sub_desc = desc + *parents, leaf = path.split(".") + for part in parents: + sub_desc = sub_desc.fields_by_name[part].message_type + field = sub_desc.fields_by_name[leaf] + # operator.attrgetter resolves dotted paths ("a.b.c") natively in C. + getter = operator.attrgetter(path) + if ( + field.type == descriptor.FieldDescriptor.TYPE_ENUM + and enum_format == "name" + ): + # Proto enum attributes return integer wire numbers; map to enum names. + names = {v.number: v.name for v in field.enum_type.values} + getter = lambda m, g=getter, n=names: n[g(m)] + getters.append(getter) + return tuple(getters) + + def to_tuple( msg: message.Message, *, @@ -208,19 +240,10 @@ def to_tuple( enum_format: Literal["name", "number"] = "name", ) -> tuple[Any, ...]: """Converts a Protobuf message instance to a tuple of field values.""" - fields = _leaf_field_paths(msg.DESCRIPTOR) if schema is None else schema - result = [] - for path in fields: - sub_msg, leaf = _resolve_leaf(msg, path) - field = sub_msg.DESCRIPTOR.fields_by_name[leaf] - val = getattr(sub_msg, leaf) - if ( - field.type == descriptor.FieldDescriptor.TYPE_ENUM - and enum_format == "name" - ): - val = field.enum_type.values_by_number[val].name - result.append(val) - return tuple(result) + # Convert schema to a tuple of field names so it is hashable for lru_cache. + paths = _leaf_field_paths(msg.DESCRIPTOR) if schema is None else tuple(schema) + getters = _compile_getters(msg.DESCRIPTOR, paths, enum_format) + return tuple(g(msg) for g in getters) def from_tuple( @@ -232,6 +255,7 @@ def from_tuple( """Populates a Protobuf message from a sequence of field values.""" desc = _resolve_descriptor(proto) fields = _leaf_field_paths(desc) if schema is None else schema + # Instantiate an empty message class from the Descriptor and populate leaves. msg = message_factory.GetMessageClass(desc)() for path, val in zip(fields, values): sub_msg, leaf = _resolve_leaf(msg, path)