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..d191045 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,161 @@ 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: + 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 + 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 _extension_declarations( + file: descriptor_pb2.FileDescriptorProto, +) -> 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_extension_declarations(message, (*package, message.name)) + +def _message_extension_declarations( + message: descriptor_pb2.DescriptorProto, + 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_extension_declarations(nested, (*path, nested.name)) -def _declares_google_protobuf_option_extension(file: descriptor_pb2.FileDescriptorProto) -> bool: - return any(extension.extendee.startswith(".google.protobuf.") for extension in file.extension) +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)) -def _declares_message_scoped_extensions(file: descriptor_pb2.FileDescriptorProto) -> bool: - return any(_message_declares_extensions(message) for message in file.message_type) + 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 _message_declares_extensions(message: descriptor_pb2.DescriptorProto) -> bool: - if message.extension: +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: + return any(type_name == root or type_name.startswith(f"{root}.") for root in helper_roots) + + +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_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( @@ -150,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/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..9770d67 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" @@ -678,10 +680,12 @@ def build_model(request: descriptor_pb2.FileDescriptorSet | object) -> Descripto 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) + 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] = {} @@ -743,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: @@ -759,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: @@ -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..bfa72d6 100644 --- a/tests/test_descriptor_set.py +++ b/tests/test_descriptor_set.py @@ -77,6 +77,110 @@ 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 _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() + 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 +218,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 +341,83 @@ 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_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( + 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 +427,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..3e81bdd --- /dev/null +++ b/tests/test_proto3_custom_option_extensions.py @@ -0,0 +1,426 @@ +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_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") + + 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 _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, + 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