Skip to content
Merged
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
52 changes: 38 additions & 14 deletions dpsynth/adapters/protobuf.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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}'."
Expand All @@ -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."
Expand All @@ -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 = []
Expand All @@ -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):
Expand All @@ -201,26 +209,41 @@ 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,
*,
schema: domain.Schema | Mapping[str, domain.AttributeType] | None = None,
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(
Expand All @@ -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)
Expand Down
Loading