Skip to content
Open
Show file tree
Hide file tree
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
1 change: 1 addition & 0 deletions .changelog/5402.fixed
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
`opentelemetry-codegen-json`: support decoding enums from names
Original file line number Diff line number Diff line change
Expand Up @@ -299,7 +299,6 @@ def _generate_imports(
std_imports = [
"builtins",
"dataclasses",
"functools",
"typing",
]
if include_enum:
Expand All @@ -310,12 +309,6 @@ def _generate_imports(

writer.blank_line()

writer.assignment(
"_dataclass",
"functools.partial(dataclasses.dataclass, slots=True)",
)
writer.blank_line()

# Collect all imports needed
imports = self._collect_imports(proto_file)
imports.add(f"import {self._get_codec_module_path()}")
Expand Down Expand Up @@ -399,7 +392,7 @@ def _generate_message_class(
msg_desc.name,
bases=(f"{codec}.JsonMessage",),
decorators=("typing.final",),
decorator_name="_dataclass",
slots=True,
):
if msg_desc.field or msg_desc.nested_type or msg_desc.enum_type:
writer.docstring([f"Generated from protobuf message {msg_desc.name}"])
Expand Down Expand Up @@ -541,7 +534,7 @@ def _generate_from_dict(
"from_dict",
["cls", "data: builtins.dict[builtins.str, typing.Any]"],
decorators=["builtins.classmethod"],
return_type=f'"{current_path}"',
return_type=f"{current_path}",
):
writer.docstring(
[
Expand Down Expand Up @@ -681,10 +674,9 @@ def _generate_deserialization_statements(
)
elif field_desc.type == descriptor.FieldDescriptorProto.TYPE_ENUM:
enum_type = self._resolve_enum_type(field_desc.type_name, proto_file)
writer.writeln(f'{codec}.validate_type({var_name}, builtins.int, "{field_desc.name}")')
writer.assignment(
f'{target_dict}["{field_desc.name}"]',
f"{enum_type}({var_name})",
f'{codec}.decode_enum({var_name}, {enum_type}, "{field_desc.name}")',
)
elif is_hex_encoded_field(field_desc.name):
writer.assignment(
Expand Down Expand Up @@ -738,7 +730,7 @@ def _get_deserialization_expr(
return f"{msg_type}.from_dict({var_name})"
if field_desc.type == descriptor.FieldDescriptorProto.TYPE_ENUM:
enum_type = self._resolve_enum_type(field_desc.type_name, proto_file)
return f"{enum_type}({var_name})"
return f'{codec}.decode_enum({var_name}, {enum_type}, "{field_desc.name}")'
if is_hex_encoded_field(field_desc.name):
return f'{codec}.decode_hex({var_name}, "{field_desc.name}")'
if is_int64_type(field_desc.type):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import abc
import base64
import collections.abc
import enum
import json
import math
import typing
Expand All @@ -14,6 +15,7 @@

T = typing.TypeVar("T")
M = typing.TypeVar("M", bound="JsonMessage")
EnumT = typing.TypeVar("EnumT", bound=enum.IntEnum)


class JsonMessage(abc.ABC):
Expand Down Expand Up @@ -237,3 +239,31 @@ def validate_type(
"""
if not isinstance(value, expected_types):
raise TypeError(f"Field '{field_name}' expected {expected_types}, got {type(value).__name__}")


def decode_enum(value: int | str, enum_type: type[EnumT], field_name: str) -> EnumT:
"""
Decode a JSON enum value into an enum member.

Per the ProtoJSON spec, parsers must accept both enum names (str)
and integer values (int).

Args:
value: The enum name or integer value to decode.
enum_type: The enum class to decode into.
field_name: The name of the field being decoded (for error messages).
Returns:
The corresponding enum member.
"""
if isinstance(value, bool):
Comment thread
herin049 marked this conversation as resolved.
raise TypeError(f"Field '{field_name}' expected int or str, got bool")
validate_type(value, (int, str), field_name)
if isinstance(value, str):
try:
return enum_type[value]
except KeyError:
raise ValueError(f"Invalid enum name '{value}' for field '{field_name}'") from None
try:
return enum_type(value)
except ValueError:
raise ValueError(f"Invalid enum value {value} for field '{field_name}'") from None
45 changes: 45 additions & 0 deletions codegen/opentelemetry-codegen-json/tests/test_json_codec.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,15 @@
# Copyright The OpenTelemetry Authors
# SPDX-License-Identifier: Apache-2.0

import enum
import math
from typing import Any

import pytest # type: ignore

from opentelemetry.codegen.json.runtime.json_codec import (
decode_base64,
decode_enum,
decode_float,
decode_hex,
decode_int64,
Expand All @@ -20,6 +23,12 @@
)


class _Color(enum.IntEnum):
RED = 0
GREEN = 1
BLUE = 2


@pytest.mark.parametrize(
"value, expected",
[
Expand Down Expand Up @@ -185,3 +194,39 @@ def test_validate_type() -> None:
match=r"Field 'field' expected \(<class 'int'>, <class 'float'>\), got str",
):
validate_type("s", (int, float), "field")


@pytest.mark.parametrize(
"value, expected",
[
(1, _Color.GREEN),
("GREEN", _Color.GREEN),
(0, _Color.RED),
("RED", _Color.RED),
],
)
def test_decode_enum(value: int | str, expected: _Color) -> None:
assert decode_enum(value, _Color, "field") is expected


@pytest.mark.parametrize(
"value, expected_error",
[
([], TypeError),
(True, TypeError),
(False, TypeError),
(99, ValueError),
("NOT_A_COLOR", ValueError),
],
)
def test_decode_enum_errors(value: Any, expected_error: type[Exception]) -> None:
with pytest.raises(expected_error):
decode_enum(value, _Color, "field")


def test_decode_enum_error_messages_include_field_name() -> None:
with pytest.raises(ValueError, match="field"):
decode_enum(99, _Color, "field")

with pytest.raises(ValueError, match="field"):
decode_enum("NOT_A_COLOR", _Color, "field")
33 changes: 32 additions & 1 deletion codegen/opentelemetry-codegen-json/tests/test_serde.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,6 +103,17 @@ def test_generated_message_roundtrip(
assert new_msg == msg


def test_enum_field_accepts_name_and_int(test_v1_types: tuple[type[Any], type[Any]]) -> None:
TestMessage, _ = test_v1_types

from_name = TestMessage.from_dict({"enumValue": "SUCCESS"})
from_int = TestMessage.from_dict({"enumValue": 1})

assert from_name.enum_value == TestMessage.TestEnum.SUCCESS
assert from_int.enum_value == TestMessage.TestEnum.SUCCESS
assert from_name == from_int


def test_cross_reference(common_v1_types: type[Any], trace_v1_types: type[Any]) -> None:
InstrumentationScope = common_v1_types
Span = trace_v1_types
Expand Down Expand Up @@ -252,6 +263,25 @@ def test_nested_enum_suite(complex_v1_types: tuple[type[Any], ...]) -> None:
assert new_msg.repeated_nested == msg.repeated_nested


def test_nested_enum_suite_accepts_names(
complex_v1_types: tuple[type[Any], ...],
) -> None:
NestedEnumSuite = complex_v1_types[3]

msg = NestedEnumSuite.from_dict(
{
"nested": "NESTED_FOO",
"repeatedNested": ["NESTED_FOO", "NESTED_BAR"],
}
)

assert msg.nested == NestedEnumSuite.NestedEnum.NESTED_FOO
assert msg.repeated_nested == [
NestedEnumSuite.NestedEnum.NESTED_FOO,
NestedEnumSuite.NestedEnum.NESTED_BAR,
]


def test_deeply_nested(complex_v1_types: tuple[type[Any], ...]) -> None:
DeeplyNested = complex_v1_types[4]

Expand Down Expand Up @@ -300,7 +330,8 @@ def test_defaults_and_none(
({"listStrings": "not a list"}, TypeError, "expected <class 'list'>"),
({"name": 123}, TypeError, "expected <class 'str'>"),
({"subMessage": "not a dict"}, TypeError, "expected <class 'dict'>"),
({"enumValue": "SUCCESS"}, TypeError, "expected <class 'int'>"),
({"enumValue": []}, TypeError, "expected"),
({"enumValue": "NOT_A_NAME"}, ValueError, None),
({"listMessages": [None]}, TypeError, "expected <class 'dict'>"),
],
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import abc
import base64
import collections.abc
import enum
import json
import math
import typing
Expand All @@ -14,6 +15,7 @@

T = typing.TypeVar("T")
M = typing.TypeVar("M", bound="JsonMessage")
EnumT = typing.TypeVar("EnumT", bound=enum.IntEnum)


class JsonMessage(abc.ABC):
Expand Down Expand Up @@ -134,9 +136,7 @@ def decode_hex(value: str | None, field_name: str) -> bytes:
try:
return bytes.fromhex(value)
except ValueError as error:
raise ValueError(
f"Invalid hex string for field '{field_name}': {error}"
) from None
raise ValueError(f"Invalid hex string for field '{field_name}': {error}") from None


def decode_base64(value: str | None, field_name: str) -> bytes:
Expand All @@ -155,9 +155,7 @@ def decode_base64(value: str | None, field_name: str) -> bytes:
try:
return base64.b64decode(value)
except Exception as error:
raise ValueError(
f"Invalid base64 string for field '{field_name}': {error}"
) from None
raise ValueError(f"Invalid base64 string for field '{field_name}': {error}") from None


def decode_int64(value: int | str | None, field_name: str) -> int:
Expand All @@ -176,9 +174,7 @@ def decode_int64(value: int | str | None, field_name: str) -> int:
try:
return int(value)
except (ValueError, TypeError):
raise ValueError(
f"Invalid int64 value for field '{field_name}': {value}"
) from None
raise ValueError(f"Invalid int64 value for field '{field_name}': {value}") from None


def decode_float(value: float | str | None, field_name: str) -> float:
Expand All @@ -203,9 +199,7 @@ def decode_float(value: float | str | None, field_name: str) -> float:
try:
return float(value)
except (ValueError, TypeError):
raise ValueError(
f"Invalid float value for field '{field_name}': {value}"
) from None
raise ValueError(f"Invalid float value for field '{field_name}': {value}") from None


def decode_repeated(
Expand Down Expand Up @@ -244,7 +238,32 @@ def validate_type(
field_name: The name of the field being validated (for error messages).
"""
if not isinstance(value, expected_types):
raise TypeError(
f"Field '{field_name}' expected {expected_types}, "
f"got {type(value).__name__}"
)
raise TypeError(f"Field '{field_name}' expected {expected_types}, got {type(value).__name__}")


def decode_enum(value: int | str, enum_type: type[EnumT], field_name: str) -> EnumT:
"""
Decode a JSON enum value into an enum member.

Per the ProtoJSON spec, parsers must accept both enum names (str)
and integer values (int).

Args:
value: The enum name or integer value to decode.
enum_type: The enum class to decode into.
field_name: The name of the field being decoded (for error messages).
Returns:
The corresponding enum member.
"""
if isinstance(value, bool):
raise TypeError(f"Field '{field_name}' expected int or str, got bool")
validate_type(value, (int, str), field_name)
if isinstance(value, str):
try:
return enum_type[value]
except KeyError:
raise ValueError(f"Invalid enum name '{value}' for field '{field_name}'") from None
try:
return enum_type(value)
except ValueError:
raise ValueError(f"Invalid enum value {value} for field '{field_name}'") from None
Original file line number Diff line number Diff line change
Expand Up @@ -8,17 +8,14 @@

import builtins
import dataclasses
import functools
import typing

_dataclass = functools.partial(dataclasses.dataclass, slots=True)

import opentelemetry.proto_json._json_codec
import opentelemetry.proto_json.logs.v1.logs


@typing.final
@_dataclass
@dataclasses.dataclass(slots=True)
class ExportLogsServiceRequest(opentelemetry.proto_json._json_codec.JsonMessage):
"""
Generated from protobuf message ExportLogsServiceRequest
Expand Down Expand Up @@ -59,7 +56,7 @@ def from_dict(cls, data: builtins.dict[builtins.str, typing.Any]) -> ExportLogsS


@typing.final
@_dataclass
@dataclasses.dataclass(slots=True)
class ExportLogsServiceResponse(opentelemetry.proto_json._json_codec.JsonMessage):
"""
Generated from protobuf message ExportLogsServiceResponse
Expand Down Expand Up @@ -100,7 +97,7 @@ def from_dict(cls, data: builtins.dict[builtins.str, typing.Any]) -> ExportLogsS


@typing.final
@_dataclass
@dataclasses.dataclass(slots=True)
class ExportLogsPartialSuccess(opentelemetry.proto_json._json_codec.JsonMessage):
"""
Generated from protobuf message ExportLogsPartialSuccess
Expand Down
Loading
Loading