From b31546cd9157942bd67357477b380a3397576f83 Mon Sep 17 00:00:00 2001 From: Anthony Printup <92564080+anthonyprintup@users.noreply.github.com> Date: Sun, 28 Jun 2026 21:40:37 +0200 Subject: [PATCH 1/2] feat: support proto3 custom option extensions Add shared classification for protobuf custom option extension declarations, including ExtensionRangeOptions. Validate proto3 extension declarations across descriptor graphs and keep unsupported extension fields out of codec generation. Refine descriptor-set discovery so pure option-definition files remain import metadata while mixed schema files are still generated. Track the new generator module in both source-tree and installed CMake package configs. --- cmake/Protocyte.cmake | 1 + cmake/protocyteConfig.cmake.in | 1 + src/protocyte/descriptor_set.py | 59 ++- src/protocyte/extensions.py | 20 + src/protocyte/model.py | 54 ++- tests/test_cmake.py | 2 + tests/test_descriptor_set.py | 78 ++++ tests/test_proto3_custom_option_extensions.py | 403 ++++++++++++++++++ 8 files changed, 603 insertions(+), 15 deletions(-) create mode 100644 src/protocyte/extensions.py create mode 100644 tests/test_proto3_custom_option_extensions.py diff --git a/cmake/Protocyte.cmake b/cmake/Protocyte.cmake index b2e36a1..f232b98 100644 --- a/cmake/Protocyte.cmake +++ b/cmake/Protocyte.cmake @@ -11,6 +11,7 @@ set( "${PROTOCYTE_PACKAGE_ROOT}/cpp.py" "${PROTOCYTE_PACKAGE_ROOT}/descriptor_set.py" "${PROTOCYTE_PACKAGE_ROOT}/errors.py" + "${PROTOCYTE_PACKAGE_ROOT}/extensions.py" "${PROTOCYTE_PACKAGE_ROOT}/main.py" "${PROTOCYTE_PACKAGE_ROOT}/model.py" "${PROTOCYTE_PACKAGE_ROOT}/parameters.py" diff --git a/cmake/protocyteConfig.cmake.in b/cmake/protocyteConfig.cmake.in index bf69be4..405cffb 100644 --- a/cmake/protocyteConfig.cmake.in +++ b/cmake/protocyteConfig.cmake.in @@ -23,6 +23,7 @@ set( "${PROTOCYTE_PACKAGE_ROOT}/cpp.py" "${PROTOCYTE_PACKAGE_ROOT}/descriptor_set.py" "${PROTOCYTE_PACKAGE_ROOT}/errors.py" + "${PROTOCYTE_PACKAGE_ROOT}/extensions.py" "${PROTOCYTE_PACKAGE_ROOT}/main.py" "${PROTOCYTE_PACKAGE_ROOT}/model.py" "${PROTOCYTE_PACKAGE_ROOT}/parameters.py" diff --git a/src/protocyte/descriptor_set.py b/src/protocyte/descriptor_set.py index 60cb0bf..55ffb93 100644 --- a/src/protocyte/descriptor_set.py +++ b/src/protocyte/descriptor_set.py @@ -8,6 +8,7 @@ from google.protobuf.message import DecodeError from protocyte.errors import ProtocyteError +from protocyte.extensions import is_custom_option_extension _RUNTIME_PREFIX = "google/protobuf/" @@ -95,26 +96,66 @@ def discover_files(descriptor_set: descriptor_pb2.FileDescriptorSet) -> list[str def _is_initial_discoverable_target(file: descriptor_pb2.FileDescriptorProto) -> bool: return ( _is_referenced_type_discoverable(file) - and not _declares_google_protobuf_option_extension(file) + and not _is_pure_custom_option_definition(file) ) def _is_referenced_type_discoverable(file: descriptor_pb2.FileDescriptorProto) -> bool: - return file.name not in _INTERNAL_DESCRIPTOR_FILES and not _declares_message_scoped_extensions(file) + return file.name not in _INTERNAL_DESCRIPTOR_FILES and not _declares_unsupported_message_scoped_extensions(file) + + +def _is_pure_custom_option_definition(file: descriptor_pb2.FileDescriptorProto) -> bool: + extensions = list(_extensions(file)) + if ( + not extensions + or any(not is_custom_option_extension(extension) for extension in extensions) + or file.service + ): + return False + helper_roots = { + _normalize_type_name(extension.type_name) + for extension in extensions + if extension.type_name + } + return all(_is_extension_helper_type(type_name, helper_roots) for type_name in _declared_type_names(file)) + + +def _extensions( + file: descriptor_pb2.FileDescriptorProto, +) -> Iterable[descriptor_pb2.FieldDescriptorProto]: + yield from file.extension + for message in file.message_type: + yield from _message_extensions(message) -def _declares_google_protobuf_option_extension(file: descriptor_pb2.FileDescriptorProto) -> bool: - return any(extension.extendee.startswith(".google.protobuf.") for extension in file.extension) +def _message_extensions( + message: descriptor_pb2.DescriptorProto, +) -> Iterable[descriptor_pb2.FieldDescriptorProto]: + yield from message.extension + for nested in message.nested_type: + yield from _message_extensions(nested) -def _declares_message_scoped_extensions(file: descriptor_pb2.FileDescriptorProto) -> bool: - return any(_message_declares_extensions(message) for message in file.message_type) +def _is_extension_helper_type(type_name: str, helper_roots: set[str]) -> bool: + return any(type_name == root or type_name.startswith(f"{root}.") for root in helper_roots) -def _message_declares_extensions(message: descriptor_pb2.DescriptorProto) -> bool: - if message.extension: +def _declares_unsupported_message_scoped_extensions(file: descriptor_pb2.FileDescriptorProto) -> bool: + return any(_message_declares_unsupported_extensions(message) for message in file.message_type) + + +def _message_declares_unsupported_extensions(message: descriptor_pb2.DescriptorProto) -> bool: + if any(not is_custom_option_extension(extension) for extension in message.extension): return True - return any(_message_declares_extensions(nested) for nested in message.nested_type) + return any(_message_declares_unsupported_extensions(nested) for nested in message.nested_type) + + +def _declared_type_names(file: descriptor_pb2.FileDescriptorProto) -> Iterable[str]: + package = tuple(part for part in file.package.split(".") if part) + for message in file.message_type: + yield from _message_type_names(package, message) + for enum in file.enum_type: + yield _fully_qualified_name((*package, enum.name)) def _index_declared_types( diff --git a/src/protocyte/extensions.py b/src/protocyte/extensions.py new file mode 100644 index 0000000..2815fe8 --- /dev/null +++ b/src/protocyte/extensions.py @@ -0,0 +1,20 @@ +from __future__ import annotations + +from google.protobuf import descriptor_pb2 + + +CUSTOM_OPTION_EXTENDEES = { + ".google.protobuf.FileOptions", + ".google.protobuf.MessageOptions", + ".google.protobuf.FieldOptions", + ".google.protobuf.OneofOptions", + ".google.protobuf.EnumOptions", + ".google.protobuf.EnumValueOptions", + ".google.protobuf.ExtensionRangeOptions", + ".google.protobuf.ServiceOptions", + ".google.protobuf.MethodOptions", +} + + +def is_custom_option_extension(field: descriptor_pb2.FieldDescriptorProto) -> bool: + return field.extendee in CUSTOM_OPTION_EXTENDEES diff --git a/src/protocyte/model.py b/src/protocyte/model.py index 2fed243..8388540 100644 --- a/src/protocyte/model.py +++ b/src/protocyte/model.py @@ -9,6 +9,7 @@ from protocyte.descriptor_set import validate_virtual_file_name from protocyte.errors import ProtocyteError +from protocyte.extensions import CUSTOM_OPTION_EXTENDEES, is_custom_option_extension FieldDescriptorProto = descriptor_pb2.FieldDescriptorProto @@ -150,6 +151,7 @@ ARRAY_OPTION_NAME = "protocyte.array" CONSTANT_OPTION_NAME = "protocyte.constant" PACKAGE_CONSTANT_OPTION_NAME = "protocyte.package_constant" +_CUSTOM_OPTION_EXTENDEES = CUSTOM_OPTION_EXTENDEES CONSTANT_KIND_BOOL = "bool" CONSTANT_KIND_INT32 = "int32" CONSTANT_KIND_INT64 = "int64" @@ -675,10 +677,13 @@ def build_model(request: descriptor_pb2.FileDescriptorSet | object) -> Descripto if missing: raise ProtocyteError(f"protoc request is missing file descriptors for: {', '.join(missing)}") + selected_files = set(file_to_generate) + for file in files_by_name.values(): + _validate_extension_declarations(file, selected_for_generation=file.name in selected_files) + for name in file_to_generate: validate_virtual_file_name(name) file = files_by_name[name] - _reject_unsupported_extension_declarations(file) _reject_unsupported_file_features(file, f"target file {name}") _validate_import_graph(files_by_name, file_to_generate) @@ -846,17 +851,52 @@ def _reject_unsupported_file_features(file: descriptor_pb2.FileDescriptorProto, raise ProtocyteError(f"{label}: protobuf Editions are not supported in v1") -def _reject_unsupported_extension_declarations(file: descriptor_pb2.FileDescriptorProto) -> None: - def reject_message_extensions(message: descriptor_pb2.DescriptorProto, path: str) -> None: - if message.extension: +def _is_custom_option_extension(field: descriptor_pb2.FieldDescriptorProto) -> bool: + return is_custom_option_extension(field) + + +def _validate_extension_declarations( + file: descriptor_pb2.FileDescriptorProto, + *, + selected_for_generation: bool, +) -> None: + syntax = _file_syntax(file) + + def validate_extension( + extension: descriptor_pb2.FieldDescriptorProto, + path: str | None, + ) -> None: + if _is_custom_option_extension(extension): + return + if syntax == "proto3": + raise ProtocyteError( + f"{file.name}: extension {_extension_full_name(file, path, extension)} " + f"extends unsupported proto3 target {extension.extendee}" + ) + if selected_for_generation and path is not None: raise ProtocyteError( f"{file.name}: message {proto_full_name(file, path)}: extension declarations are not supported" ) + + def validate_message_extensions(message: descriptor_pb2.DescriptorProto, path: str) -> None: + for extension in message.extension: + validate_extension(extension, path) for nested in message.nested_type: - reject_message_extensions(nested, f"{path}.{nested.name}") + validate_message_extensions(nested, f"{path}.{nested.name}") + for extension in file.extension: + validate_extension(extension, None) for message in file.message_type: - reject_message_extensions(message, message.name) + validate_message_extensions(message, message.name) + + +def _extension_full_name( + file: descriptor_pb2.FileDescriptorProto, + path: str | None, + extension: descriptor_pb2.FieldDescriptorProto, +) -> str: + extension_path = f"{path}.{extension.name}" if path is not None else extension.name + return proto_full_name(file, extension_path) def _build_enum( @@ -1312,6 +1352,8 @@ def _build_field( enums: dict[str, EnumModel], custom_options: _CustomOptions, ) -> FieldModel: + if proto.extendee: + raise ProtocyteError(f"{owner.full_name}.{proto.name}: extension fields are not supported for codec generation") if proto.type == FieldDescriptorProto.TYPE_GROUP: raise ProtocyteError(f"{owner.full_name}.{proto.name}: groups are not supported") diff --git a/tests/test_cmake.py b/tests/test_cmake.py index 4b1fdc7..9a3e6c3 100644 --- a/tests/test_cmake.py +++ b/tests/test_cmake.py @@ -34,6 +34,8 @@ def test_installed_cmake_config_tracks_descriptor_set_helper() -> None: assert '"${PROTOCYTE_PACKAGE_ROOT}/descriptor_set.py"' in source_config assert '"${PROTOCYTE_PACKAGE_ROOT}/descriptor_set.py"' in installed_config + assert '"${PROTOCYTE_PACKAGE_ROOT}/extensions.py"' in source_config + assert '"${PROTOCYTE_PACKAGE_ROOT}/extensions.py"' in installed_config def test_posix_wrapper_shell_quotes_single_quotes(tmp_path: Path) -> None: diff --git a/tests/test_descriptor_set.py b/tests/test_descriptor_set.py index 27b976c..a60dca0 100644 --- a/tests/test_descriptor_set.py +++ b/tests/test_descriptor_set.py @@ -77,6 +77,33 @@ def _custom_options_file() -> descriptor_pb2.FileDescriptorProto: return file +def _extension_range_options_file() -> descriptor_pb2.FileDescriptorProto: + file = descriptor_pb2.FileDescriptorProto() + file.name = "custom/extension_range_options.proto" + file.package = "custom" + file.syntax = "proto3" + file.dependency.append("google/protobuf/descriptor.proto") + extension = file.extension.add() + extension.name = "range_label" + extension.number = 50001 + extension.label = descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL + extension.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + extension.extendee = ".google.protobuf.ExtensionRangeOptions" + return file + + +def _mixed_custom_options_file() -> descriptor_pb2.FileDescriptorProto: + file = _custom_options_file() + message = file.message_type.add() + message.name = "PublicPayload" + field = message.field.add() + field.name = "id" + field.number = 1 + field.label = descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL + field.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + return file + + def _file_with_nested_extension() -> descriptor_pb2.FileDescriptorProto: file = descriptor_pb2.FileDescriptorProto() file.name = "custom/nested_options.proto" @@ -114,6 +141,21 @@ def _file_with_top_level_extension() -> descriptor_pb2.FileDescriptorProto: return file +def _proto3_file_with_google_protobuf_non_option_extension() -> descriptor_pb2.FileDescriptorProto: + file = descriptor_pb2.FileDescriptorProto() + file.name = "custom/timestamp_extension.proto" + file.package = "custom" + file.syntax = "proto3" + file.dependency.append("google/protobuf/timestamp.proto") + extension = file.extension.add() + extension.name = "timestamp_marker" + extension.number = 50000 + extension.label = descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL + extension.type = descriptor_pb2.FieldDescriptorProto.TYPE_INT32 + extension.extendee = ".google.protobuf.Timestamp" + return file + + def _file_with_custom_marker_field(name: str) -> descriptor_pb2.FileDescriptorProto: file = _file(name, "custom/options.proto") field = file.message_type[0].field.add() @@ -222,6 +264,29 @@ def test_discover_files_skips_imported_custom_option_extension_descriptors(tmp_p assert discover_files(load_descriptor_set(path)) == ["api/request.proto"] +def test_discover_files_skips_imported_extension_range_option_descriptors(tmp_path: Path) -> None: + path = tmp_path / "descriptor_set.pb" + _write_descriptor_set( + path, + _file("google/protobuf/descriptor.proto"), + _extension_range_options_file(), + _file("api/request.proto", "custom/extension_range_options.proto"), + ) + + assert discover_files(load_descriptor_set(path)) == ["api/request.proto"] + + +def test_discover_files_includes_custom_option_files_with_public_messages(tmp_path: Path) -> None: + path = tmp_path / "descriptor_set.pb" + _write_descriptor_set( + path, + _file("google/protobuf/descriptor.proto"), + _mixed_custom_options_file(), + ) + + assert discover_files(load_descriptor_set(path)) == ["custom/options.proto"] + + def test_discover_files_includes_user_files_with_top_level_extension_declarations( tmp_path: Path, ) -> None: @@ -231,6 +296,19 @@ def test_discover_files_includes_user_files_with_top_level_extension_declaration assert discover_files(load_descriptor_set(path)) == ["legacy.proto"] +def test_discover_files_includes_non_option_google_protobuf_extension_descriptors( + tmp_path: Path, +) -> None: + path = tmp_path / "descriptor_set.pb" + _write_descriptor_set( + path, + _timestamp_file(), + _proto3_file_with_google_protobuf_non_option_extension(), + ) + + assert discover_files(load_descriptor_set(path)) == ["custom/timestamp_extension.proto"] + + def test_discover_files_includes_extension_descriptors_referenced_by_message_fields(tmp_path: Path) -> None: path = tmp_path / "descriptor_set.pb" _write_descriptor_set( diff --git a/tests/test_proto3_custom_option_extensions.py b/tests/test_proto3_custom_option_extensions.py new file mode 100644 index 0000000..9af91e5 --- /dev/null +++ b/tests/test_proto3_custom_option_extensions.py @@ -0,0 +1,403 @@ +from types import SimpleNamespace + +import pytest +from google.protobuf import descriptor_pb2, descriptor_pool, message_factory +from google.protobuf.compiler import plugin_pb2 + +from protocyte.descriptor_set import load_descriptor_set, validate_descriptor_set +from protocyte.errors import ProtocyteError +from protocyte.model import build_model +from protocyte.plugin import generate_response + + +F = descriptor_pb2.FieldDescriptorProto + + +@pytest.mark.parametrize( + "extendee", + [ + ".google.protobuf.FileOptions", + ".google.protobuf.MessageOptions", + ".google.protobuf.FieldOptions", + ".google.protobuf.OneofOptions", + ".google.protobuf.EnumOptions", + ".google.protobuf.EnumValueOptions", + ".google.protobuf.ExtensionRangeOptions", + ".google.protobuf.ServiceOptions", + ".google.protobuf.MethodOptions", + ], +) +def test_accepts_each_supported_proto3_custom_option_extendee(extendee: str) -> None: + file = _proto3_file_with_single_custom_option_extension(extendee) + + build_model(_request(file, selected=["single_option.proto"])) + + +def test_accepts_proto3_custom_option_extensions_selected_for_generation() -> None: + file = _custom_options_file() + + model = build_model(_request(file, selected=["example/options.proto"])) + + assert model.files["example/options.proto"].messages[0].full_name == "example.options.AccessPolicy" + assert model.files["example/options.proto"].enums[0].full_name == "example.options.RpcKind" + + response = generate_response(_plugin_request(file, selected=["example/options.proto"])) + + assert not response.error + files = {item.name: item.content for item in response.file} + assert "example/options.protocyte.hpp" in files + assert "service_name" not in files["example/options.protocyte.hpp"] + assert "method_kind" not in files["example/options.protocyte.hpp"] + assert "access_policy" not in files["example/options.protocyte.hpp"] + + +def test_rejects_proto3_non_option_top_level_extensions() -> None: + file = _proto3_file_with_top_level_extension(".example.options.AccessPolicy") + + with pytest.raises( + ProtocyteError, + match=( + r"example/options\.proto: extension example\.options\.ordinary_extension " + r"extends unsupported proto3 target \.example\.options\.AccessPolicy" + ), + ): + build_model(_request(file, selected=["example/options.proto"])) + + +def test_rejects_unselected_proto3_non_option_extension_dependency() -> None: + options_file = _proto3_file_with_top_level_extension(".example.options.AccessPolicy") + consumer_file = _consumer_file_without_custom_options() + + with pytest.raises( + ProtocyteError, + match=( + r"example/options\.proto: extension example\.options\.ordinary_extension " + r"extends unsupported proto3 target \.example\.options\.AccessPolicy" + ), + ): + build_model(_request(options_file, consumer_file, selected=["example/api.proto"])) + + +def test_rejects_proto3_nested_non_option_extensions() -> None: + file = _proto3_file_with_nested_extension(".example.options.AccessPolicy") + + with pytest.raises( + ProtocyteError, + match=( + r"nested_options\.proto: extension example\.options\.Holder\.nested_extension " + r"extends unsupported proto3 target \.example\.options\.AccessPolicy" + ), + ): + build_model(_request(file, selected=["nested_options.proto"])) + + +def test_accepts_proto3_nested_custom_option_extensions() -> None: + file = _proto3_file_with_nested_extension(".google.protobuf.MethodOptions") + + model = build_model(_request(file, selected=["nested_options.proto"])) + + assert model.files["nested_options.proto"].messages[0].full_name == "example.options.Holder" + + +def test_descriptor_set_selected_option_definition_file_generates_helpers( + tmp_path, +) -> None: + path = tmp_path / "descriptor_set.pb" + _write_descriptor_set( + path, + descriptor_pb2.FileDescriptorProto.FromString(descriptor_pb2.DESCRIPTOR.serialized_pb), + _custom_options_file(), + ) + loaded = load_descriptor_set(path) + validate_descriptor_set(loaded, ["example/options.proto"]) + request = plugin_pb2.CodeGeneratorRequest() + request.file_to_generate.append("example/options.proto") + request.proto_file.extend(loaded.file) + + response = generate_response(request) + + assert not response.error + assert {item.name for item in response.file} == { + "example/options.protocyte.cpp", + "example/options.protocyte.hpp", + } + + +def test_descriptor_set_custom_options_remain_decodable_from_input( + tmp_path, +) -> None: + options_file = _custom_options_file() + consumer_file = _consumer_file_with_custom_options(options_file) + path = tmp_path / "descriptor_set.pb" + _write_descriptor_set( + path, + descriptor_pb2.FileDescriptorProto.FromString(descriptor_pb2.DESCRIPTOR.serialized_pb), + options_file, + consumer_file, + ) + loaded = load_descriptor_set(path) + validate_descriptor_set(loaded, ["example/api.proto"]) + request = plugin_pb2.CodeGeneratorRequest() + request.file_to_generate.append("example/api.proto") + request.proto_file.extend(loaded.file) + + response = generate_response(request) + + assert not response.error + decoded = _decode_consumer_custom_options(options_file, consumer_file) + assert decoded == { + "file": "api-file", + "message": "request-message", + "field": "request-id", + "service": "example", + "method_kind": 0, + "method_roles": ["user", "admin"], + } + + +def _request( + *files: descriptor_pb2.FileDescriptorProto, selected: list[str] +) -> SimpleNamespace: + return SimpleNamespace(proto_file=list(files), file_to_generate=selected) + + +def _plugin_request( + *files: descriptor_pb2.FileDescriptorProto, selected: list[str] +) -> plugin_pb2.CodeGeneratorRequest: + request = plugin_pb2.CodeGeneratorRequest() + request.file_to_generate.extend(selected) + request.proto_file.extend(files) + return request + + +def _write_descriptor_set(path, *files: descriptor_pb2.FileDescriptorProto) -> None: + descriptor_set = descriptor_pb2.FileDescriptorSet() + descriptor_set.file.extend(files) + path.write_bytes(descriptor_set.SerializeToString()) + + +def _custom_options_file() -> descriptor_pb2.FileDescriptorProto: + file = descriptor_pb2.FileDescriptorProto() + file.name = "example/options.proto" + file.package = "example.options" + file.syntax = "proto3" + file.dependency.append("google/protobuf/descriptor.proto") + + access_policy = file.message_type.add() + access_policy.name = "AccessPolicy" + roles = access_policy.field.add() + roles.name = "roles" + roles.number = 1 + roles.label = F.LABEL_REPEATED + roles.type = F.TYPE_STRING + + enum = file.enum_type.add() + enum.name = "RpcKind" + value = enum.value.add() + value.name = "RPC_KIND_REQUEST_RESPONSE" + value.number = 0 + value = enum.value.add() + value.name = "RPC_KIND_EVENT" + value.number = 1 + + _add_extension(file.extension.add(), "file_label", 50000, F.TYPE_STRING, ".google.protobuf.FileOptions") + _add_extension( + file.extension.add(), "message_label", 50001, F.TYPE_STRING, ".google.protobuf.MessageOptions" + ) + _add_extension(file.extension.add(), "field_label", 50002, F.TYPE_STRING, ".google.protobuf.FieldOptions") + _add_extension(file.extension.add(), "service_name", 50003, F.TYPE_STRING, ".google.protobuf.ServiceOptions") + _add_extension( + file.extension.add(), + "method_kind", + 50004, + F.TYPE_ENUM, + ".google.protobuf.MethodOptions", + type_name=".example.options.RpcKind", + ) + _add_extension( + file.extension.add(), + "access_policy", + 50005, + F.TYPE_MESSAGE, + ".google.protobuf.MethodOptions", + type_name=".example.options.AccessPolicy", + ) + return file + + +def _proto3_file_with_top_level_extension(extendee: str) -> descriptor_pb2.FileDescriptorProto: + file = _custom_options_file() + extension = file.extension.add() + _add_extension(extension, "ordinary_extension", 50100, F.TYPE_INT32, extendee) + return file + + +def _proto3_file_with_nested_extension(extendee: str) -> descriptor_pb2.FileDescriptorProto: + file = descriptor_pb2.FileDescriptorProto() + file.name = "nested_options.proto" + file.package = "example.options" + file.syntax = "proto3" + file.dependency.append("google/protobuf/descriptor.proto") + message = file.message_type.add() + message.name = "Holder" + extension = message.extension.add() + _add_extension(extension, "nested_extension", 50000, F.TYPE_STRING, extendee) + return file + + +def _proto3_file_with_single_custom_option_extension( + extendee: str, +) -> descriptor_pb2.FileDescriptorProto: + file = descriptor_pb2.FileDescriptorProto() + file.name = "single_option.proto" + file.package = "example.options" + file.syntax = "proto3" + file.dependency.append("google/protobuf/descriptor.proto") + extension = file.extension.add() + _add_extension(extension, "custom_option", 50000, F.TYPE_STRING, extendee) + return file + + +def _consumer_file_with_custom_options( + options_file: descriptor_pb2.FileDescriptorProto, +) -> descriptor_pb2.FileDescriptorProto: + file = descriptor_pb2.FileDescriptorProto() + file.name = "example/api.proto" + file.package = "example.api" + file.syntax = "proto3" + file.dependency.append("example/options.proto") + file.options.ParseFromString(_option_bytes(options_file, "FileOptions", "file_label", "api-file")) + + message = file.message_type.add() + message.name = "Request" + message.options.ParseFromString( + _option_bytes(options_file, "MessageOptions", "message_label", "request-message") + ) + field = message.field.add() + field.name = "id" + field.number = 1 + field.label = F.LABEL_OPTIONAL + field.type = F.TYPE_STRING + field.options.ParseFromString(_option_bytes(options_file, "FieldOptions", "field_label", "request-id")) + + response = file.message_type.add() + response.name = "Response" + + service = file.service.add() + service.name = "ExampleService" + service.options.ParseFromString( + _option_bytes(options_file, "ServiceOptions", "service_name", "example") + ) + method = service.method.add() + method.name = "Ping" + method.input_type = ".example.api.Request" + method.output_type = ".example.api.Response" + method.options.ParseFromString(_method_options_bytes(options_file)) + return file + + +def _consumer_file_without_custom_options() -> descriptor_pb2.FileDescriptorProto: + file = descriptor_pb2.FileDescriptorProto() + file.name = "example/api.proto" + file.package = "example.api" + file.syntax = "proto3" + file.dependency.append("example/options.proto") + message = file.message_type.add() + message.name = "Request" + field = message.field.add() + field.name = "id" + field.number = 1 + field.label = F.LABEL_OPTIONAL + field.type = F.TYPE_STRING + return file + + +def _add_extension( + extension: descriptor_pb2.FieldDescriptorProto, + name: str, + number: int, + field_type: int, + extendee: str, + *, + type_name: str = "", +) -> None: + extension.name = name + extension.number = number + extension.label = F.LABEL_OPTIONAL + extension.type = field_type + extension.extendee = extendee + if type_name: + extension.type_name = type_name + + +def _option_bytes( + options_file: descriptor_pb2.FileDescriptorProto, + option_type: str, + extension_name: str, + value: str, +) -> bytes: + pool = _pool_with_options(options_file) + options_desc = pool.FindMessageTypeByName(f"google.protobuf.{option_type}") + options_cls = message_factory.GetMessageClass(options_desc) + extension = pool.FindExtensionByName(f"example.options.{extension_name}") + options = options_cls() + options.Extensions[extension] = value + return options.SerializeToString() + + +def _method_options_bytes(options_file: descriptor_pb2.FileDescriptorProto) -> bytes: + pool = _pool_with_options(options_file) + options_desc = pool.FindMessageTypeByName("google.protobuf.MethodOptions") + options_cls = message_factory.GetMessageClass(options_desc) + method_kind = pool.FindExtensionByName("example.options.method_kind") + access_policy = pool.FindExtensionByName("example.options.access_policy") + options = options_cls() + options.Extensions[method_kind] = 0 + options.Extensions[access_policy].roles.extend(["user", "admin"]) + return options.SerializeToString() + + +def _decode_consumer_custom_options( + options_file: descriptor_pb2.FileDescriptorProto, + consumer_file: descriptor_pb2.FileDescriptorProto, +) -> dict[str, object]: + pool = _pool_with_options(options_file) + pool.Add(consumer_file) + file_options_cls = message_factory.GetMessageClass(pool.FindMessageTypeByName("google.protobuf.FileOptions")) + message_options_cls = message_factory.GetMessageClass(pool.FindMessageTypeByName("google.protobuf.MessageOptions")) + field_options_cls = message_factory.GetMessageClass(pool.FindMessageTypeByName("google.protobuf.FieldOptions")) + service_options_cls = message_factory.GetMessageClass(pool.FindMessageTypeByName("google.protobuf.ServiceOptions")) + method_options_cls = message_factory.GetMessageClass(pool.FindMessageTypeByName("google.protobuf.MethodOptions")) + + file_options = file_options_cls() + file_options.ParseFromString(consumer_file.options.SerializeToString()) + message_options = message_options_cls() + message_options.ParseFromString(consumer_file.message_type[0].options.SerializeToString()) + field_options = field_options_cls() + field_options.ParseFromString(consumer_file.message_type[0].field[0].options.SerializeToString()) + service_options = service_options_cls() + service_options.ParseFromString(consumer_file.service[0].options.SerializeToString()) + method_options = method_options_cls() + method_options.ParseFromString(consumer_file.service[0].method[0].options.SerializeToString()) + + file_label = pool.FindExtensionByName("example.options.file_label") + message_label = pool.FindExtensionByName("example.options.message_label") + field_label = pool.FindExtensionByName("example.options.field_label") + service_name = pool.FindExtensionByName("example.options.service_name") + method_kind = pool.FindExtensionByName("example.options.method_kind") + access_policy = pool.FindExtensionByName("example.options.access_policy") + return { + "file": file_options.Extensions[file_label], + "message": message_options.Extensions[message_label], + "field": field_options.Extensions[field_label], + "service": service_options.Extensions[service_name], + "method_kind": method_options.Extensions[method_kind], + "method_roles": list(method_options.Extensions[access_policy].roles), + } + + +def _pool_with_options(options_file: descriptor_pb2.FileDescriptorProto) -> descriptor_pool.DescriptorPool: + pool = descriptor_pool.DescriptorPool() + pool.AddSerializedFile(descriptor_pb2.DESCRIPTOR.serialized_pb) + pool.Add(options_file) + return pool From afc01cb5a22c00dd20953c56d1c861de8f02268f Mon Sep 17 00:00:00 2001 From: Anthony Printup <92564080+anthonyprintup@users.noreply.github.com> Date: Sun, 28 Jun 2026 22:07:37 +0200 Subject: [PATCH 2/2] fix: address proto3 option review feedback Scope proto3 extension validation to the selected import closure so unrelated descriptor-set files cannot block generation. Teach descriptor-set discovery to skip transitive custom-option helper types and namespace-only nested extension containers while still selecting option files that expose public fields or nested messages. Add regression coverage for transitive helper enums, nested scalar option namespaces, public namespace guards, and unrelated invalid proto3 extensions outside the selected graph. --- src/protocyte/descriptor_set.py | 135 +++++++++++++++--- src/protocyte/model.py | 12 +- tests/test_descriptor_set.py | 131 +++++++++++++++++ tests/test_proto3_custom_option_extensions.py | 23 +++ 4 files changed, 278 insertions(+), 23 deletions(-) diff --git a/src/protocyte/descriptor_set.py b/src/protocyte/descriptor_set.py index 55ffb93..d191045 100644 --- a/src/protocyte/descriptor_set.py +++ b/src/protocyte/descriptor_set.py @@ -105,35 +105,109 @@ def _is_referenced_type_discoverable(file: descriptor_pb2.FileDescriptorProto) - def _is_pure_custom_option_definition(file: descriptor_pb2.FileDescriptorProto) -> bool: - extensions = list(_extensions(file)) + declarations = list(_extension_declarations(file)) + extensions = [extension for extension, _ in declarations] if ( not extensions or any(not is_custom_option_extension(extension) for extension in extensions) or file.service ): return False - helper_roots = { - _normalize_type_name(extension.type_name) - for extension in extensions - if extension.type_name - } - return all(_is_extension_helper_type(type_name, helper_roots) for type_name in _declared_type_names(file)) + declared_type_names = set(_declared_type_names(file)) + helper_roots = _custom_option_helper_roots(file, declarations, declared_type_names) + return all(_is_extension_helper_type(type_name, helper_roots) for type_name in declared_type_names) -def _extensions( +def _extension_declarations( file: descriptor_pb2.FileDescriptorProto, -) -> Iterable[descriptor_pb2.FieldDescriptorProto]: - yield from file.extension +) -> Iterable[tuple[descriptor_pb2.FieldDescriptorProto, str | None]]: + for extension in file.extension: + yield extension, None + package = tuple(part for part in file.package.split(".") if part) for message in file.message_type: - yield from _message_extensions(message) + yield from _message_extension_declarations(message, (*package, message.name)) -def _message_extensions( +def _message_extension_declarations( message: descriptor_pb2.DescriptorProto, -) -> Iterable[descriptor_pb2.FieldDescriptorProto]: - yield from message.extension + path: tuple[str, ...], +) -> Iterable[tuple[descriptor_pb2.FieldDescriptorProto, str]]: + scope = _fully_qualified_name(path) + for extension in message.extension: + yield extension, scope for nested in message.nested_type: - yield from _message_extensions(nested) + yield from _message_extension_declarations(nested, (*path, nested.name)) + + +def _custom_option_helper_roots( + file: descriptor_pb2.FileDescriptorProto, + declarations: Iterable[tuple[descriptor_pb2.FieldDescriptorProto, str | None]], + declared_type_names: set[str], +) -> set[str]: + helper_roots = set[str]() + for extension, _ in declarations: + if extension.type_name: + helper_roots.add(_normalize_type_name(extension.type_name)) + + declared_messages = _index_declared_message_types(file) + _expand_custom_option_helper_roots(declared_messages, declared_type_names, helper_roots) + _add_namespace_container_roots(declared_messages, helper_roots) + return helper_roots + + +def _expand_custom_option_helper_roots( + declared_messages: dict[str, descriptor_pb2.DescriptorProto], + declared_type_names: set[str], + helper_roots: set[str], +) -> None: + stack = list(helper_roots) + while stack: + root = stack.pop() + for type_name, message in declared_messages.items(): + if not _is_extension_helper_type(type_name, {root}): + continue + for referenced in _message_direct_referenced_type_names(message): + normalized = _normalize_type_name(referenced) + if normalized in declared_type_names and not _is_extension_helper_type(normalized, helper_roots): + helper_roots.add(normalized) + stack.append(normalized) + + +def _add_namespace_container_roots( + declared_messages: dict[str, descriptor_pb2.DescriptorProto], + helper_roots: set[str], +) -> None: + namespace_roots = set[str]() + changed = True + while changed: + changed = False + for type_name, message in declared_messages.items(): + if type_name in namespace_roots or _is_extension_helper_type(type_name, helper_roots): + continue + if _is_namespace_only_custom_option_scope(type_name, message, helper_roots, namespace_roots): + namespace_roots.add(type_name) + changed = True + helper_roots.update(namespace_roots) + + +def _is_namespace_only_custom_option_scope( + type_name: str, + message: descriptor_pb2.DescriptorProto, + helper_roots: set[str], + namespace_roots: set[str], +) -> bool: + if message.field or message.oneof_decl or message.extension_range or message.reserved_range or message.reserved_name: + return False + child_type_names = [ + *(f"{type_name}.{nested.name}" for nested in message.nested_type), + *(f"{type_name}.{enum.name}" for enum in message.enum_type), + ] + if any( + not _is_extension_helper_type(child_type_name, helper_roots) and child_type_name not in namespace_roots + for child_type_name in child_type_names + ): + return False + return bool(message.extension or child_type_names) def _is_extension_helper_type(type_name: str, helper_roots: set[str]) -> bool: @@ -158,6 +232,27 @@ def _declared_type_names(file: descriptor_pb2.FileDescriptorProto) -> Iterable[s yield _fully_qualified_name((*package, enum.name)) +def _index_declared_message_types( + file: descriptor_pb2.FileDescriptorProto, +) -> dict[str, descriptor_pb2.DescriptorProto]: + declared: dict[str, descriptor_pb2.DescriptorProto] = {} + package = tuple(part for part in file.package.split(".") if part) + for message in file.message_type: + _index_message_type(declared, package, message) + return declared + + +def _index_message_type( + declared: dict[str, descriptor_pb2.DescriptorProto], + prefix: tuple[str, ...], + message: descriptor_pb2.DescriptorProto, +) -> None: + path = (*prefix, message.name) + declared[_fully_qualified_name(path)] = message + for nested in message.nested_type: + _index_message_type(declared, path, nested) + + def _index_declared_types( files: Iterable[descriptor_pb2.FileDescriptorProto], ) -> dict[str, str]: @@ -191,12 +286,18 @@ def _referenced_type_names(file: descriptor_pb2.FileDescriptorProto) -> Iterable def _message_referenced_type_names( message: descriptor_pb2.DescriptorProto, +) -> Iterable[str]: + yield from _message_direct_referenced_type_names(message) + for nested in message.nested_type: + yield from _message_referenced_type_names(nested) + + +def _message_direct_referenced_type_names( + message: descriptor_pb2.DescriptorProto, ) -> Iterable[str]: for field in message.field: if field.type_name: yield field.type_name - for nested in message.nested_type: - yield from _message_referenced_type_names(nested) def _fully_qualified_name(parts: Iterable[str]) -> str: diff --git a/src/protocyte/model.py b/src/protocyte/model.py index 8388540..9770d67 100644 --- a/src/protocyte/model.py +++ b/src/protocyte/model.py @@ -677,16 +677,15 @@ def build_model(request: descriptor_pb2.FileDescriptorSet | object) -> Descripto if missing: raise ProtocyteError(f"protoc request is missing file descriptors for: {', '.join(missing)}") - selected_files = set(file_to_generate) - for file in files_by_name.values(): - _validate_extension_declarations(file, selected_for_generation=file.name in selected_files) - for name in file_to_generate: validate_virtual_file_name(name) file = files_by_name[name] _reject_unsupported_file_features(file, f"target file {name}") - _validate_import_graph(files_by_name, file_to_generate) + selected_files = set(file_to_generate) + reachable_files = _validate_import_graph(files_by_name, file_to_generate) + for name in reachable_files: + _validate_extension_declarations(files_by_name[name], selected_for_generation=name in selected_files) files: dict[str, FileModel] = {} messages: dict[str, MessageModel] = {} @@ -748,7 +747,7 @@ def _index_request_files( def _validate_import_graph( files: dict[str, descriptor_pb2.FileDescriptorProto], roots: Iterable[str], -) -> None: +) -> set[str]: stack = list(roots) seen: set[str] = set() while stack: @@ -764,6 +763,7 @@ def _validate_import_graph( if dependency not in files: raise ProtocyteError(f"{name} imports missing descriptor {dependency}") stack.append(dependency) + return seen def _custom_options(proto_files: Iterable[descriptor_pb2.FileDescriptorProto]) -> _CustomOptions: diff --git a/tests/test_descriptor_set.py b/tests/test_descriptor_set.py index a60dca0..bfa72d6 100644 --- a/tests/test_descriptor_set.py +++ b/tests/test_descriptor_set.py @@ -92,6 +92,83 @@ def _extension_range_options_file() -> descriptor_pb2.FileDescriptorProto: return file +def _custom_options_file_with_transitive_helper_enum() -> descriptor_pb2.FileDescriptorProto: + file = descriptor_pb2.FileDescriptorProto() + file.name = "custom/policy_options.proto" + file.package = "custom" + file.syntax = "proto3" + file.dependency.append("google/protobuf/descriptor.proto") + + policy = file.message_type.add() + policy.name = "Policy" + level = policy.field.add() + level.name = "level" + level.number = 1 + level.label = descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL + level.type = descriptor_pb2.FieldDescriptorProto.TYPE_ENUM + level.type_name = ".custom.Severity" + + severity = file.enum_type.add() + severity.name = "Severity" + value = severity.value.add() + value.name = "SEVERITY_UNSPECIFIED" + value.number = 0 + value = severity.value.add() + value.name = "SEVERITY_HIGH" + value.number = 1 + + extension = file.extension.add() + extension.name = "policy" + extension.number = 50000 + extension.label = descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL + extension.type = descriptor_pb2.FieldDescriptorProto.TYPE_MESSAGE + extension.type_name = ".custom.Policy" + extension.extendee = ".google.protobuf.MethodOptions" + return file + + +def _nested_namespace_custom_options_file() -> descriptor_pb2.FileDescriptorProto: + file = descriptor_pb2.FileDescriptorProto() + file.name = "custom/nested_method_options.proto" + file.package = "custom" + file.syntax = "proto3" + file.dependency.append("google/protobuf/descriptor.proto") + + namespace = file.message_type.add() + namespace.name = "Opts" + extension = namespace.extension.add() + extension.name = "tag" + extension.number = 50000 + extension.label = descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL + extension.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + extension.extendee = ".google.protobuf.MethodOptions" + return file + + +def _nested_namespace_custom_options_file_with_public_field() -> descriptor_pb2.FileDescriptorProto: + file = _nested_namespace_custom_options_file() + file.name = "custom/nested_public_field_options.proto" + field = file.message_type[0].field.add() + field.name = "public_id" + field.number = 1 + field.label = descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL + field.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + return file + + +def _nested_namespace_custom_options_file_with_public_nested_message() -> descriptor_pb2.FileDescriptorProto: + file = _nested_namespace_custom_options_file() + file.name = "custom/nested_public_message_options.proto" + message = file.message_type[0].nested_type.add() + message.name = "PublicPayload" + field = message.field.add() + field.name = "id" + field.number = 1 + field.label = descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL + field.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + return file + + def _mixed_custom_options_file() -> descriptor_pb2.FileDescriptorProto: file = _custom_options_file() message = file.message_type.add() @@ -276,6 +353,60 @@ def test_discover_files_skips_imported_extension_range_option_descriptors(tmp_pa assert discover_files(load_descriptor_set(path)) == ["api/request.proto"] +def test_discover_files_skips_transitive_custom_option_helper_types(tmp_path: Path) -> None: + path = tmp_path / "descriptor_set.pb" + _write_descriptor_set( + path, + _file("google/protobuf/descriptor.proto"), + _custom_options_file_with_transitive_helper_enum(), + _file("api/request.proto", "custom/policy_options.proto"), + ) + + assert discover_files(load_descriptor_set(path)) == ["api/request.proto"] + + +def test_discover_files_skips_nested_scalar_custom_option_namespaces(tmp_path: Path) -> None: + path = tmp_path / "descriptor_set.pb" + _write_descriptor_set( + path, + _file("google/protobuf/descriptor.proto"), + _nested_namespace_custom_options_file(), + _file("api/request.proto", "custom/nested_method_options.proto"), + ) + + assert discover_files(load_descriptor_set(path)) == ["api/request.proto"] + + +def test_discover_files_includes_nested_custom_option_namespaces_with_public_fields(tmp_path: Path) -> None: + path = tmp_path / "descriptor_set.pb" + _write_descriptor_set( + path, + _file("google/protobuf/descriptor.proto"), + _nested_namespace_custom_options_file_with_public_field(), + _file("api/request.proto", "custom/nested_public_field_options.proto"), + ) + + assert discover_files(load_descriptor_set(path)) == [ + "api/request.proto", + "custom/nested_public_field_options.proto", + ] + + +def test_discover_files_includes_nested_custom_option_namespaces_with_public_nested_messages(tmp_path: Path) -> None: + path = tmp_path / "descriptor_set.pb" + _write_descriptor_set( + path, + _file("google/protobuf/descriptor.proto"), + _nested_namespace_custom_options_file_with_public_nested_message(), + _file("api/request.proto", "custom/nested_public_message_options.proto"), + ) + + assert discover_files(load_descriptor_set(path)) == [ + "api/request.proto", + "custom/nested_public_message_options.proto", + ] + + def test_discover_files_includes_custom_option_files_with_public_messages(tmp_path: Path) -> None: path = tmp_path / "descriptor_set.pb" _write_descriptor_set( diff --git a/tests/test_proto3_custom_option_extensions.py b/tests/test_proto3_custom_option_extensions.py index 9af91e5..3e81bdd 100644 --- a/tests/test_proto3_custom_option_extensions.py +++ b/tests/test_proto3_custom_option_extensions.py @@ -78,6 +78,14 @@ def test_rejects_unselected_proto3_non_option_extension_dependency() -> None: build_model(_request(options_file, consumer_file, selected=["example/api.proto"])) +def test_ignores_unrelated_proto3_non_option_extension_outside_selected_graph() -> None: + invalid_file = _proto3_file_with_top_level_extension(".example.options.AccessPolicy") + invalid_file.name = "unrelated/options.proto" + consumer_file = _consumer_file_with_no_imports() + + build_model(_request(invalid_file, consumer_file, selected=["example/api.proto"])) + + def test_rejects_proto3_nested_non_option_extensions() -> None: file = _proto3_file_with_nested_extension(".example.options.AccessPolicy") @@ -312,6 +320,21 @@ def _consumer_file_without_custom_options() -> descriptor_pb2.FileDescriptorProt return file +def _consumer_file_with_no_imports() -> descriptor_pb2.FileDescriptorProto: + file = descriptor_pb2.FileDescriptorProto() + file.name = "example/api.proto" + file.package = "example.api" + file.syntax = "proto3" + message = file.message_type.add() + message.name = "Request" + field = message.field.add() + field.name = "id" + field.number = 1 + field.label = F.LABEL_OPTIONAL + field.type = F.TYPE_STRING + return file + + def _add_extension( extension: descriptor_pb2.FieldDescriptorProto, name: str,