use std::collections::HashMap;
use prost_types::{field_descriptor_proto, DescriptorProto};
use super::sparse_field_map::{SparseFieldMap, MAX_INLINE_CAPACITY};
const LABEL_REPEATED: i32 = 3;
const ROOT_PREFIX: &str = "";
#[derive(Clone, Debug)]
#[repr(C)]
pub struct FieldInfo {
pub storage_index: usize,
pub oneof_index: Option<i32>,
pub field_type: field_descriptor_proto::Type,
pub is_scalar: bool,
pub is_repeated: bool,
pub name: Box<str>,
pub type_name: Option<String>,
}
#[derive(Clone, Debug)]
pub struct OneofMember {
pub is_scalar: bool,
pub storage_index: usize,
}
#[derive(Debug)]
pub struct DescriptorWithFieldCache {
fields: SparseFieldMap<FieldInfo>,
pub is_map_entry: bool,
pub scalar_count: usize,
pub complex_count: usize,
oneof_groups: Vec<Vec<OneofMember>>,
}
impl DescriptorWithFieldCache {
#[inline(always)]
pub fn from_descriptor(desc: &DescriptorProto) -> Self {
let max_inline_field_num = desc
.field
.iter()
.filter_map(|f| f.number)
.filter(|&n| n < MAX_INLINE_CAPACITY as i32)
.max()
.unwrap_or(0);
let mut fields = SparseFieldMap::new(max_inline_field_num);
let mut scalar_count = 0usize;
let mut complex_count = 0usize;
let oneof_count = desc.oneof_decl.len();
let mut oneof_groups: Vec<Vec<OneofMember>> = vec![Vec::new(); oneof_count];
for field in desc.field.iter() {
if let Some(num) = field.number {
let field_type = field.r#type();
let is_repeated = field.label == Some(LABEL_REPEATED);
let is_scalar = !is_repeated && field_type != field_descriptor_proto::Type::Message;
let storage_index = if is_scalar {
let idx = scalar_count;
scalar_count += 1;
idx
} else {
let idx = complex_count;
complex_count += 1;
idx
};
let oneof_index = field
.oneof_index
.filter(|_| field.proto3_optional != Some(true));
if let Some(idx) = oneof_index {
if let Some(group) = oneof_groups.get_mut(idx as usize) {
group.push(OneofMember {
is_scalar,
storage_index,
});
}
}
let info = FieldInfo {
name: Box::from(field.name.as_deref().unwrap_or("")),
field_type,
type_name: field.type_name.clone(),
is_repeated,
is_scalar,
storage_index,
oneof_index,
};
fields.insert(num, info);
}
}
let is_map_entry = desc
.options
.as_ref()
.and_then(|o| o.map_entry)
.unwrap_or(false);
DescriptorWithFieldCache {
fields,
is_map_entry,
scalar_count,
complex_count,
oneof_groups,
}
}
#[inline(always)]
pub fn get_field(&self, field_num: i32) -> Option<&FieldInfo> {
self.fields.get(field_num)
}
#[inline(always)]
pub fn get_oneof_group(&self, oneof_index: i32) -> &[OneofMember] {
self.oneof_groups
.get(oneof_index as usize)
.map(|v| v.as_slice())
.unwrap_or(&[])
}
#[inline(always)]
pub fn oneof_groups(&self) -> &[Vec<OneofMember>] {
&self.oneof_groups
}
}
#[derive(Debug)]
pub struct MessageRegistry {
messages: HashMap<String, DescriptorWithFieldCache>,
pub(crate) root_descriptor: DescriptorWithFieldCache,
}
impl MessageRegistry {
#[inline(always)]
pub fn from_descriptor(root: &DescriptorProto) -> Self {
let mut messages = HashMap::new();
Self::collect_messages(root, &mut messages);
let root_descriptor = DescriptorWithFieldCache::from_descriptor(root);
MessageRegistry {
messages,
root_descriptor,
}
}
#[inline(always)]
fn collect_messages(
root: &DescriptorProto,
acc: &mut HashMap<String, DescriptorWithFieldCache>,
) {
let mut stack: Vec<(&DescriptorProto, String)> = vec![(root, ROOT_PREFIX.to_string())];
while let Some((desc, current_prefix)) = stack.pop() {
let name = desc.name.as_deref().unwrap_or("");
let full_name = if current_prefix.is_empty() {
format!(".{name}")
} else {
format!("{current_prefix}.{name}")
};
acc.insert(
full_name.clone(),
DescriptorWithFieldCache::from_descriptor(desc),
);
for nested in desc.nested_type.iter().rev() {
stack.push((nested, full_name.clone()));
}
}
}
pub fn get(&self, type_name: &str) -> Option<&DescriptorWithFieldCache> {
self.messages.get(type_name)
}
pub fn get_field_type_name(&self, field_num: i32) -> Option<&str> {
self.root_descriptor
.get_field(field_num)?
.type_name
.as_deref()
}
}
#[cfg(test)]
pub mod tests {
use super::*;
use crate::zeroparser::parser::tests::{make_descriptor, make_field};
impl MessageRegistry {
pub fn get_field_name(&self, field_num: i32) -> Option<&str> {
self.root_descriptor.get_field(field_num).map(|f| &*f.name)
}
}
#[test]
fn descriptor_cache_field_lookup() {
let fields = vec![
make_field(1, "id", field_descriptor_proto::Type::Int32, false, None),
make_field(2, "name", field_descriptor_proto::Type::String, false, None),
make_field(
200,
"large",
field_descriptor_proto::Type::String,
false,
None,
),
make_field(3, "items", field_descriptor_proto::Type::Int32, true, None),
];
let desc = make_descriptor("TestMessage", fields);
let cache = DescriptorWithFieldCache::from_descriptor(&desc);
let field1 = cache.get_field(1).unwrap();
assert_eq!(&*field1.name, "id");
assert_eq!(field1.field_type, field_descriptor_proto::Type::Int32);
assert!(!field1.is_repeated);
let field2 = cache.get_field(2).unwrap();
assert_eq!(&*field2.name, "name");
let field200 = cache.get_field(200).unwrap();
assert_eq!(&*field200.name, "large");
let field3 = cache.get_field(3).unwrap();
assert!(field3.is_repeated);
assert!(cache.get_field(99).is_none());
assert!(cache.get_field(0).is_none());
}
#[test]
fn descriptor_cache_map_entry() {
let desc = make_descriptor("TestMessage", vec![]);
assert!(!DescriptorWithFieldCache::from_descriptor(&desc).is_map_entry);
let mut map_desc = make_descriptor(
"MapEntry",
vec![
make_field(1, "key", field_descriptor_proto::Type::String, false, None),
make_field(2, "value", field_descriptor_proto::Type::Int32, false, None),
],
);
map_desc.options = Some(prost_types::MessageOptions {
map_entry: Some(true),
..Default::default()
});
assert!(DescriptorWithFieldCache::from_descriptor(&map_desc).is_map_entry);
}
#[test]
fn message_registry_lookup() {
let fields = vec![make_field(
1,
"id",
field_descriptor_proto::Type::Int32,
false,
None,
)];
let desc = make_descriptor("RootMessage", fields);
let registry = MessageRegistry::from_descriptor(&desc);
assert_eq!(&*registry.root_descriptor.get_field(1).unwrap().name, "id");
assert!(registry.get(".RootMessage").is_some());
assert!(registry.get(".NonExistent").is_none());
}
#[test]
fn message_registry_nested() {
let level3 = make_descriptor(
"Level3",
vec![make_field(
1,
"field3",
field_descriptor_proto::Type::Bool,
false,
None,
)],
);
let mut level2 = make_descriptor("Level2", vec![]);
level2.nested_type.push(level3);
let mut level1 = make_descriptor(
"Level1",
vec![
make_field(1, "id", field_descriptor_proto::Type::Int32, false, None),
make_field(
2,
"nested",
field_descriptor_proto::Type::Message,
false,
Some(".Level1.Level2"),
),
],
);
level1.nested_type.push(level2);
let registry = MessageRegistry::from_descriptor(&level1);
assert!(registry.get(".Level1").is_some());
assert!(registry.get(".Level1.Level2").is_some());
assert!(registry.get(".Level1.Level2.Level3").is_some());
let level3_cache = registry.get(".Level1.Level2.Level3").unwrap();
assert_eq!(&*level3_cache.get_field(1).unwrap().name, "field3");
assert_eq!(registry.get_field_type_name(1), None);
assert_eq!(registry.get_field_type_name(2), Some(".Level1.Level2"));
assert_eq!(registry.get_field_type_name(99), None);
}
}