Skip to content
Merged
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 cmake/Protocyte.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
1 change: 1 addition & 0 deletions cmake/protocyteConfig.cmake.in
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
164 changes: 153 additions & 11 deletions src/protocyte/descriptor_set.py
Original file line number Diff line number Diff line change
Expand Up @@ -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/"
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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:
Expand Down
20 changes: 20 additions & 0 deletions src/protocyte/extensions.py
Original file line number Diff line number Diff line change
@@ -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
58 changes: 50 additions & 8 deletions src/protocyte/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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] = {}
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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")

Expand Down
2 changes: 2 additions & 0 deletions tests/test_cmake.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Loading
Loading