Skip to content

Commit 16e6da7

Browse files
committed
feat: OpenEnum, Config::typed_enum_fields
Introduce OpenEnum to represent fields of Protobuf enumerated types with the possibility of unknown values. A new enum_type annotation is supported in the prost attribute inside derives, which allows to specify the type-checked representation of enum types in message fields and oneof variants. The accepted values are "open" or "closed". Use OpenEnum in generated code for enum-typed fields of messages instead of i32, if the `enum_type` annotation is set to "open". OpenEnum features get and set methods, which are meant to be the primary ways to get and set open enumeration fields, instead of the getters and setters generated by Message derive. Add typed_enum_fields method to prost-build configuration, which allows type-checked representation of enumerations in fields of message structs and variants of oneof enums. The argument and the invocation order works like with the boxed method. Depending on the syntax (and preparing for the future support of editions), the type-checked representation can be closed (for proto2) or open (for proto3). The former is represented by the generated enum type itself, while the latter is represented by OpenEnum wrapping the enum type.
1 parent 1e93f56 commit 16e6da7

16 files changed

Lines changed: 26540 additions & 414 deletions

File tree

‎benchmarks/benches/dataset.rs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,7 +37,7 @@ fn benchmark_dataset<M>(criterion: &mut Criterion, name: &str, dataset: &'static
3737
where
3838
M: prost::Message + Default + 'static,
3939
{
40-
let mut group = criterion.benchmark_group(&format!("dataset/{}", name));
40+
let mut group = criterion.benchmark_group(format!("dataset/{}", name));
4141

4242
group.bench_function("merge", move |b| {
4343
let dataset = load_dataset(dataset).unwrap();

‎prost-build/src/code_generator.rs‎

Lines changed: 65 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -388,7 +388,7 @@ impl<'b> CodeGenerator<'_, 'b> {
388388

389389
fn append_field(&mut self, fq_message_name: &str, field: &Field) {
390390
let type_ = field.descriptor.r#type();
391-
let repeated = field.descriptor.label() == Label::Repeated;
391+
let repeated = field.descriptor.label == Some(Label::Repeated as i32);
392392
let deprecated = self.deprecated(&field.descriptor);
393393
let optional = self.optional(&field.descriptor);
394394
let boxed = self
@@ -415,12 +415,16 @@ impl<'b> CodeGenerator<'_, 'b> {
415415
let type_tag = self.field_type_tag(&field.descriptor);
416416
self.buf.push_str(&type_tag);
417417

418-
if type_ == Type::Bytes {
419-
let bytes_type = self
420-
.context
421-
.bytes_type(fq_message_name, field.descriptor.name());
422-
self.buf
423-
.push_str(&format!("={:?}", bytes_type.annotation()));
418+
match type_ {
419+
Type::Bytes => {
420+
let bytes_type = self
421+
.context
422+
.bytes_type(fq_message_name, field.descriptor.name());
423+
self.buf
424+
.push_str(&format!("={:?}", bytes_type.annotation()));
425+
}
426+
Type::Enum => self.push_enum_type_annotation(fq_message_name, field.descriptor.name()),
427+
_ => {}
424428
}
425429

426430
match field.descriptor.label() {
@@ -537,12 +541,16 @@ impl<'b> CodeGenerator<'_, 'b> {
537541
let value_tag = self.map_value_type_tag(value);
538542

539543
self.buf.push_str(&format!(
540-
"#[prost({}=\"{}, {}\", tag=\"{}\")]\n",
544+
"#[prost({}=\"{}, {}\"",
541545
map_type.annotation(),
542546
key_tag,
543547
value_tag,
544-
field.descriptor.number()
545548
));
549+
if value.r#type() == Type::Enum {
550+
self.push_enum_type_annotation(fq_message_name, field.descriptor.name());
551+
}
552+
self.buf
553+
.push_str(&format!(", tag=\"{}\")]\n", field.descriptor.number()));
546554
self.append_field_attributes(fq_message_name, field.descriptor.name());
547555
self.push_indent();
548556
self.buf.push_str(&format!(
@@ -630,11 +638,12 @@ impl<'b> CodeGenerator<'_, 'b> {
630638

631639
self.push_indent();
632640
let ty_tag = self.field_type_tag(&field.descriptor);
633-
self.buf.push_str(&format!(
634-
"#[prost({}, tag=\"{}\")]\n",
635-
ty_tag,
636-
field.descriptor.number()
637-
));
641+
self.buf.push_str(&format!("#[prost({}", ty_tag,));
642+
if field.descriptor.r#type() == Type::Enum {
643+
self.push_enum_type_annotation(&oneof_name, field.descriptor.name());
644+
}
645+
self.buf
646+
.push_str(&format!(", tag=\"{}\")]\n", field.descriptor.number()));
638647
self.append_field_attributes(&oneof_name, field.descriptor.name());
639648

640649
self.push_indent();
@@ -930,13 +939,21 @@ impl<'b> CodeGenerator<'_, 'b> {
930939
self.buf.push_str("}\n");
931940
}
932941

942+
fn push_enum_type_annotation(&mut self, fq_message_name: &str, field_name: &str) {
943+
match self.enum_field_repr(fq_message_name, field_name) {
944+
EnumRepr::Int => {}
945+
EnumRepr::Open => self.buf.push_str(", enum_type=\"open\""),
946+
EnumRepr::Closed => self.buf.push_str(", enum_type=\"closed\""),
947+
}
948+
}
949+
933950
fn resolve_type(&self, field: &FieldDescriptorProto, fq_message_name: &str) -> String {
934951
match field.r#type() {
935952
Type::Float => String::from("f32"),
936953
Type::Double => String::from("f64"),
937954
Type::Uint32 | Type::Fixed32 => String::from("u32"),
938955
Type::Uint64 | Type::Fixed64 => String::from("u64"),
939-
Type::Int32 | Type::Sfixed32 | Type::Sint32 | Type::Enum => String::from("i32"),
956+
Type::Int32 | Type::Sfixed32 | Type::Sint32 => String::from("i32"),
940957
Type::Int64 | Type::Sfixed64 | Type::Sint64 => String::from("i64"),
941958
Type::Bool => String::from("bool"),
942959
Type::String => format!("{}::alloc::string::String", self.context.prost_path()),
@@ -946,6 +963,15 @@ impl<'b> CodeGenerator<'_, 'b> {
946963
.rust_type()
947964
.to_owned(),
948965
Type::Group | Type::Message => self.resolve_ident(field.type_name()),
966+
Type::Enum => match self.enum_field_repr(fq_message_name, field.name()) {
967+
EnumRepr::Int => String::from("i32"),
968+
EnumRepr::Open => format!(
969+
"{}::OpenEnum<{}>",
970+
self.context.prost_path(),
971+
self.resolve_ident(field.type_name())
972+
),
973+
EnumRepr::Closed => self.resolve_ident(field.type_name()),
974+
},
949975
}
950976
}
951977

@@ -987,6 +1013,24 @@ impl<'b> CodeGenerator<'_, 'b> {
9871013
.join("::")
9881014
}
9891015

1016+
fn enum_field_repr(&self, fq_message_name: &str, field_name: &str) -> EnumRepr {
1017+
if self
1018+
.context
1019+
.is_typed_enum_field(fq_message_name, field_name)
1020+
{
1021+
// FIXME: store information for the code generator to know when
1022+
// proto3 enums are used in proto2, where they should be open
1023+
// accordingly to the spec:
1024+
// https://protobuf.dev/programming-guides/enum/#spec
1025+
match self.syntax {
1026+
Syntax::Proto2 => EnumRepr::Closed,
1027+
Syntax::Proto3 => EnumRepr::Open,
1028+
}
1029+
} else {
1030+
EnumRepr::Int
1031+
}
1032+
}
1033+
9901034
fn field_type_tag(&self, field: &FieldDescriptorProto) -> Cow<'static, str> {
9911035
match field.r#type() {
9921036
Type::Float => Cow::Borrowed("float"),
@@ -1077,6 +1121,12 @@ fn can_pack(field: &FieldDescriptorProto) -> bool {
10771121
)
10781122
}
10791123

1124+
enum EnumRepr {
1125+
Int,
1126+
Closed,
1127+
Open,
1128+
}
1129+
10801130
struct EnumVariantMapping<'a> {
10811131
path_idx: usize,
10821132
proto_name: &'a str,

‎prost-build/src/config.rs‎

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@ pub struct Config {
3737
pub(crate) enum_attributes: PathMap<String>,
3838
pub(crate) field_attributes: PathMap<String>,
3939
pub(crate) boxed: PathMap<()>,
40+
pub(crate) typed_enum_fields: PathMap<()>,
4041
pub(crate) prost_types: bool,
4142
pub(crate) strip_enum_prefix: bool,
4243
pub(crate) out_dir: Option<PathBuf>,
@@ -374,6 +375,30 @@ impl Config {
374375
self
375376
}
376377

378+
/// Represent Protobuf enum types encountered in matched fields with types
379+
/// bound to their corresponding Rust enum types, rather than the default `i32`.
380+
///
381+
/// Depending on the proto file syntax, the representation type can be:
382+
/// * For closed enums (in proto2), the corresponding Rust enum type.
383+
/// * For open enums (in proto3), the Rust enum type wrapped in [`OpenEnum`](prost::OpenEnum).
384+
///
385+
/// # Arguments
386+
///
387+
/// **`path`** - a path matching any number of fields. These fields will get the type-checked
388+
/// enum representation.
389+
/// For details about matching fields see [`btree_map`](#method.btree_map).
390+
///
391+
/// # Examples
392+
///
393+
/// ```rust
394+
/// # let mut config = prost_build::Config::new();
395+
/// config.typed_enum_fields(".my_messages");
396+
/// ```
397+
pub fn typed_enum_fields(&mut self, path: impl AsRef<str>) -> &mut Self {
398+
self.typed_enum_fields.insert(path.as_ref().to_owned(), ());
399+
self
400+
}
401+
377402
/// Configures the code generator to use the provided service generator.
378403
pub fn service_generator(&mut self, service_generator: Box<dyn ServiceGenerator>) -> &mut Self {
379404
self.service_generator = Some(service_generator);
@@ -1179,6 +1204,7 @@ impl default::Default for Config {
11791204
enum_attributes: PathMap::default(),
11801205
field_attributes: PathMap::default(),
11811206
boxed: PathMap::default(),
1207+
typed_enum_fields: PathMap::default(),
11821208
prost_types: true,
11831209
strip_enum_prefix: true,
11841210
out_dir: None,

‎prost-build/src/context.rs‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -168,6 +168,15 @@ impl<'a> Context<'a> {
168168
false
169169
}
170170

171+
/// Returns `true` if the named message field should be generated using the
172+
/// type-safe `OpenEnum` representation.
173+
pub fn is_typed_enum_field(&self, fq_message_name: &str, field_name: &str) -> bool {
174+
self.config
175+
.typed_enum_fields
176+
.get_first_field(fq_message_name, field_name)
177+
.is_some()
178+
}
179+
171180
/// Returns `true` if this message can automatically derive Copy trait.
172181
pub fn can_message_derive_copy(&self, fq_message_name: &str) -> bool {
173182
assert_eq!(".", &fq_message_name[..1]);

0 commit comments

Comments
 (0)