use crate::attrs::{bilrost_attrs, TagList};
use crate::crate_name;
use crate::field::traits::{DecodeLifetime, DecodeMode, FieldBearer, Tagged, WhereFor};
use alloc::boxed::Box;
use alloc::collections::{BTreeMap, BTreeSet};
use alloc::format;
use alloc::string::{String, ToString};
use alloc::vec;
use alloc::vec::Vec;
use core::iter::repeat;
use core::mem::take;
use core::ops::Deref;
use core::{iter, slice};
use eyre::{bail, eyre as err, Report as Error};
use itertools::{repeat_n, Either, Itertools};
use proc_macro2::{Ident, TokenStream};
use quote::{quote, ToTokens};
use syn::{parse_str, Attribute, Type};
mod ignored;
mod oneof;
pub mod traits;
mod value;
pub use ignored::{initializer_class_definition, InitMode};
pub use value::OneofVariant;
#[derive(Clone)]
pub struct Field {
ident: TokenStream,
content: MessageFieldContent,
}
#[derive(Clone)]
enum MessageFieldContent {
Value(Box<value::MessageField>),
Oneof(Box<oneof::OneofInclusion>),
Ignored(Box<ignored::IgnoredField>),
}
use MessageFieldContent::*;
fn ident_string(ident: &impl ToString) -> String {
let s = ident.to_string();
if let Some(after_hash) = s.strip_prefix("r#") {
after_hash.to_string()
} else {
s
}
}
pub fn parse_message_fields(
fields: syn::Fields,
init_mode: InitMode,
reserved: Option<TagList>,
) -> Result<Vec<Field>, Error> {
let mut next_tag = Some(match fields {
syn::Fields::Unnamed(..) => 0,
_ => 1,
});
let unsorted_fields = fields
.iter()
.enumerate()
.map(|(index, field)| {
let field_ident = field
.ident
.as_ref()
.map(ToTokens::to_token_stream)
.unwrap_or_else(|| {
let index = syn::Index::from(index);
quote!(#index)
});
let field = Field::new(
&field_ident,
&field.ty,
&field.attrs,
next_tag,
init_mode.clone(),
)
.map_err(|e| err!("invalid field {field_ident}: {e}"))?;
if !field.is_ignored() {
next_tag = field.last_tag().checked_add(1);
}
Ok(field)
})
.collect::<Result<Vec<_>, Error>>()?;
let all_tags: BTreeMap<u32, &Field> = unsorted_fields
.iter()
.flat_map(|field| field.tags().into_iter().zip(repeat(field)))
.collect();
for reserved_range in reserved.unwrap_or_default().iter_tag_ranges() {
if let Some((forbidden_tag, bad_field)) = all_tags.range(reserved_range).next() {
let field_ident = bad_field.ident();
bail!("field {field_ident} has reserved tag {forbidden_tag}");
}
}
if let Some((duplicated_tag, _)) = unsorted_fields
.iter()
.flat_map(|field| field.tags())
.sorted_unstable()
.tuple_windows()
.find(|(a, b)| a == b)
{
bail!("multiple fields have tag {duplicated_tag}")
};
Ok(unsorted_fields)
}
pub fn tag_measurer<T: Tagged>(for_these: impl IntoIterator<Item = T>) -> TokenStream {
let crate_ = crate_name();
if matches!(for_these.into_iter().flat_map(|t| t.tags()).max(), Some(max_tag) if max_tag < 32) {
quote!(#crate_::encoding::TrivialTagMeasurer)
} else {
quote!(#crate_::encoding::RuntimeTagMeasurer)
}
}
impl Field {
pub fn new(
ident: &TokenStream,
ty: &Type,
attrs: &[Attribute],
inferred_tag: Option<u32>,
init_mode: InitMode,
) -> Result<Field, Error> {
let attrs = bilrost_attrs(attrs)?;
Ok(Field {
content: if let Some(field) = ignored::IgnoredField::new(ty, &attrs, init_mode)? {
Ignored(field)
} else if let Some(field) = oneof::OneofInclusion::new(ty, &attrs)? {
Oneof(field)
} else {
Value(value::MessageField::new(ty, attrs, inferred_tag)?)
},
ident: ident.clone(),
})
}
pub fn ident(&self) -> &TokenStream {
&self.ident
}
pub fn ty(&self) -> &Type {
match &self.content {
Value(value) => value.ty(),
Oneof(oneof) => oneof.ty(),
Ignored(_) => panic!("ignored fields have no type"),
}
}
pub fn is_ignored(&self) -> bool {
matches!(self.content, Ignored(..))
}
pub fn ignored_and_uses_struct_update_syntax(&self) -> bool {
matches!(&self.content, Ignored(inner) if inner.uses_struct_update_syntax())
}
pub fn has_enumeration_type(&self) -> bool {
let Value(scalar) = &self.content else {
return false;
};
scalar.has_enumeration_type()
}
pub fn tag_list_guard(&self) -> Option<TokenStream> {
let Oneof(field) = &self.content else {
return None; };
let crate_ = crate_name();
let mut tags = self.tags();
tags.sort();
let oneof_ty = &field.ty;
let description = format!(
"tags don't match for oneof field {field_name} with type {oneof_ty_name}",
field_name = self.ident,
oneof_ty_name = oneof_ty.to_token_stream(),
);
let description = description.as_str();
Some(quote!(
#crate_::assert_tags_are_equal(
#description,
<#oneof_ty as #crate_::encoding::Oneof>::FIELD_TAGS,
&[#(#tags),*],
);
))
}
pub fn encode(&self, instance: &FieldTarget) -> TokenStream {
let target = instance.const_field_ref(self);
match &self.content {
Value(scalar) => scalar.encode(target),
Oneof(oneof) => oneof.encode(target),
Ignored(..) => panic!("cannot encode ignored field"),
}
}
pub fn prepend(&self, instance: &FieldTarget) -> TokenStream {
let target = instance.const_field_ref(self);
match &self.content {
Value(scalar) => scalar.prepend(target),
Oneof(oneof) => oneof.prepend(target),
Ignored(..) => panic!("cannot prepend ignored field"),
}
}
pub fn decode(
&self,
instance: &FieldTarget,
lifetime: DecodeLifetime,
mode: DecodeMode,
) -> TokenStream {
let target = instance.mut_field_ref(self);
match &self.content {
Value(scalar) => scalar.decode(target, lifetime, mode),
Oneof(oneof) => oneof.decode(target, lifetime, mode),
Ignored(..) => panic!("canot decode ignored field"),
}
}
pub fn encoded_len(&self, instance: &FieldTarget) -> TokenStream {
let target = instance.const_field_ref(self);
match &self.content {
Value(scalar) => scalar.encoded_len(target),
Oneof(oneof) => oneof.encoded_len(target),
Ignored(..) => panic!("cannot get encode length of ignored field"),
}
}
pub fn empty(&self, variant_tag: Option<u32>) -> Option<TokenStream> {
match &self.content {
Value(scalar) => Some(scalar.empty()),
Oneof(oneof) => Some(oneof.empty()),
Ignored(ignored) => ignored.initialize(variant_tag, &self.ident),
}
}
pub fn is_empty(&self, instance: &FieldTarget) -> TokenStream {
let target = instance.const_field_ref(self);
match &self.content {
Value(scalar) => scalar.is_empty(target),
Oneof(oneof) => oneof.is_empty(target),
Ignored(..) => panic!("cannot detect empty on ignored field"),
}
}
pub fn clear(&self, instance: &FieldTarget) -> TokenStream {
let target = instance.mut_field_ref(self);
match &self.content {
Value(scalar) => scalar.clear(target),
Oneof(oneof) => oneof.clear(target),
Ignored(..) => panic!("cannot clear ignored field"),
}
}
pub fn current_tag(&self, instance: &FieldTarget) -> TokenStream {
let Oneof(field) = &self.content else {
panic!("tried to use a value field as a oneof")
};
let target = instance.const_field_ref(self);
field.current_tag(target)
}
pub fn methods(&self) -> Option<TokenStream> {
match &self.content {
Value(scalar) => scalar.methods(&self.ident),
_ => None,
}
}
pub fn initializer_method(&self, variant_tag: Option<u32>) -> Option<TokenStream> {
if let Ignored(ignored) = &self.content {
ignored.initializer_method(variant_tag, &self.ident)
} else {
None
}
}
}
impl FieldBearer for Field {
fn where_terms(&self, purpose: WhereFor) -> Vec<TokenStream> {
match &self.content {
Value(field) => field.where_terms(purpose),
Oneof(field) => field.where_terms(purpose),
Ignored(field) => field.where_terms(purpose),
}
}
}
impl Tagged for Field {
fn tags(&self) -> Vec<u32> {
match &self.content {
Value(scalar) => scalar.tags(),
Oneof(oneof) => oneof.tags(),
Ignored(..) => vec![],
}
}
}
impl Tagged for &Field {
fn tags(&self) -> Vec<u32> {
(**self).tags()
}
}
struct MustMove<T>(Option<T>);
impl<T> MustMove<T> {
fn new(t: T) -> Self {
Self(Some(t))
}
fn into_inner(mut self) -> T {
take(&mut self.0).expect("MustMove value was moved out twice")
}
}
impl<T> Drop for MustMove<T> {
fn drop(&mut self) {
if self.0.is_some() {
panic!("a must-use value was dropped!");
}
}
}
impl<T> Deref for MustMove<T> {
type Target = T;
fn deref(&self) -> &T {
self.0
.as_ref()
.expect("MustMove dereferenced after the value was moved out")
}
}
pub struct MessageFieldsSorted<'a> {
chunks: Vec<FieldChunk<'a>>,
tag_measurer_ty: TokenStream,
}
enum FieldChunk<'a> {
AlwaysOrdered(&'a Field),
SortGroup(Vec<SortGroupPart<'a>>),
}
use FieldChunk::*;
enum SortGroupPart<'a> {
Contiguous(Vec<&'a Field>),
OneofPart(&'a Field),
}
use SortGroupPart::*;
struct SortGroupConfig<FC, FO> {
direction: Direction,
contiguous_part_fn: FC,
oneof_part_fn: FO,
part_fn_ty: TokenStream,
invoke_parts: TokenStream,
}
type ReversibleFields<'a> =
Either<slice::Iter<'a, &'a Field>, iter::Rev<slice::Iter<'a, &'a Field>>>;
#[derive(Copy, Clone)]
enum Direction {
Forward,
Reverse,
}
impl Direction {
fn align<T, I>(self, iterable: T) -> Either<I, iter::Rev<I>>
where
T: IntoIterator<IntoIter = I>,
I: DoubleEndedIterator,
{
match self {
Direction::Forward => Either::Left(iterable.into_iter()),
Direction::Reverse => Either::Right(iterable.into_iter().rev()),
}
}
}
fn process_sort_groups<FC, FO>(
parts: &[SortGroupPart],
instance: &FieldTarget,
config: SortGroupConfig<FC, FO>,
) -> TokenStream
where
FC: Fn(ReversibleFields) -> TokenStream,
FO: Fn(&Field) -> TokenStream,
{
let SortGroupConfig {
direction,
contiguous_part_fn,
oneof_part_fn,
part_fn_ty,
invoke_parts,
} = config;
let guaranteed_parts: Vec<_> = direction
.align(parts)
.flat_map(|part| match part {
Contiguous(fields) => {
let Some(first_field) = fields.first() else {
panic!("empty contiguous field group");
};
let first_tag = first_field.first_tag();
let closure = contiguous_part_fn(direction.align(fields));
Some(quote! { (#first_tag, ::core::option::Option::Some(#closure)) })
}
_ => None,
})
.collect();
let populate_oneof_parts: Vec<_> = direction
.align(parts)
.flat_map(|part| match part {
OneofPart(field) => {
let current_tag = field.current_tag(instance);
let closure = oneof_part_fn(field);
Some(quote! {
if let ::core::option::Option::Some(tag) = #current_tag {
parts[nparts] = (tag, ::core::option::Option::Some(#closure));
nparts += 1;
}
})
}
_ => None,
})
.collect();
let filler_none = quote!((0u32, ::core::option::Option::None));
let non_guaranteed_filler = repeat_n(&filler_none, populate_oneof_parts.len());
let num_guaranteed_parts = guaranteed_parts.len();
let max_parts = parts.len();
let sort_tag_expr = match direction {
Direction::Forward => quote!(*tag),
Direction::Reverse => quote!(::core::cmp::Reverse(*tag)),
};
let prelude = instance.prelude_for(
parts
.iter()
.flat_map(|part| match part {
Contiguous(fields) => fields.as_slice(),
OneofPart(field) => slice::from_ref(field),
})
.cloned(),
);
quote! {
{
#prelude
let mut parts: [(u32, ::core::option::Option<#part_fn_ty>); #max_parts] = [
#(#guaranteed_parts,)*
#(#non_guaranteed_filler,)*
];
let mut nparts = #num_guaranteed_parts;
#(#populate_oneof_parts)*
let parts = &mut parts[..nparts];
<[_]>::sort_unstable_by_key(parts, |(tag, _)| #sort_tag_expr);
#invoke_parts
}
}
}
impl<'a> MessageFieldsSorted<'a> {
pub fn new(unsorted_fields: impl IntoIterator<Item = &'a Field>) -> Self {
let mut chunks: Vec<FieldChunk> = vec![];
let mut fields = unsorted_fields
.into_iter()
.sorted_unstable_by_key(|field| field.first_tag())
.peekable();
let mut current_contiguous_group: Vec<&Field> = vec![];
let mut current_sort_group: Vec<SortGroupPart> = vec![];
let mut sort_group_oneof_tags = BTreeSet::<u32>::new();
while let (Some(this_field), next_field) = (fields.next(), fields.peek()) {
let this_field = MustMove::new(this_field);
let field = this_field.deref();
let first_tag = field.first_tag();
let last_tag = field.last_tag();
let overlaps =
matches!(next_field, Some(next_field) if last_tag > next_field.first_tag());
let in_current_sort_group =
matches!(sort_group_oneof_tags.iter().next_back(), Some(&end) if end > first_tag);
if in_current_sort_group {
if overlaps {
if !current_contiguous_group.is_empty() {
current_sort_group.push(Contiguous(take(&mut current_contiguous_group)));
}
sort_group_oneof_tags.extend(field.tags());
current_sort_group.push(OneofPart(this_field.into_inner()));
} else if sort_group_oneof_tags
.range(first_tag..=last_tag)
.next()
.is_some()
{
if !current_contiguous_group.is_empty() {
current_sort_group.push(Contiguous(take(&mut current_contiguous_group)));
}
current_sort_group.push(OneofPart(this_field.into_inner()));
} else {
if let Some(previous_field) = current_contiguous_group.last() {
if sort_group_oneof_tags
.range(previous_field.last_tag()..=first_tag)
.next()
.is_some()
{
current_sort_group
.push(Contiguous(take(&mut current_contiguous_group)));
}
}
current_contiguous_group.push(this_field.into_inner());
}
} else {
if overlaps {
sort_group_oneof_tags.clear();
sort_group_oneof_tags.extend(field.tags());
current_sort_group.push(OneofPart(this_field.into_inner()));
} else {
chunks.push(AlwaysOrdered(this_field.into_inner()));
}
}
if let Some(&sort_group_end) = sort_group_oneof_tags.iter().next_back() {
if !matches!(
next_field,
Some(next_field) if next_field.first_tag() < sort_group_end
) {
if !current_contiguous_group.is_empty() {
current_sort_group.push(Contiguous(take(&mut current_contiguous_group)));
}
assert!(
!current_sort_group.is_empty(),
"emitting a sort group but there are no fields"
);
chunks.push(SortGroup(take(&mut current_sort_group)));
sort_group_oneof_tags.clear();
}
}
}
assert!(
current_sort_group.into_iter().next().is_none(),
"fields left over after chunking"
);
assert!(
current_contiguous_group.into_iter().next().is_none(),
"fields left over after chunking"
);
Self {
chunks,
tag_measurer_ty: tag_measurer(fields),
}
}
pub fn new_filtering_ignored(unsorted_fields: impl IntoIterator<Item = &'a Field>) -> Self {
Self::new(
unsorted_fields
.into_iter()
.filter(|field| !field.is_ignored()),
)
}
pub fn encoded_len(&self, instance: &FieldTarget) -> TokenStream {
let tag_measurer_ty = &self.tag_measurer_ty;
let sort_group_instance;
let part_fn_instance_ty;
if instance.has_instance() {
sort_group_instance = instance.clone();
part_fn_instance_ty = quote!(Self);
} else {
sort_group_instance = FieldTarget::RefsInstance(quote!(refs));
part_fn_instance_ty = quote!(__BilrostRefs);
}
let sort_group_self = sort_group_instance.self_expr();
let renamed = sort_group_instance.rename(quote!(instance));
let chunks = self.chunks.iter().map(|chunk| match chunk {
AlwaysOrdered(field) => field.encoded_len(instance),
SortGroup(parts) => process_sort_groups(
parts,
&sort_group_instance,
SortGroupConfig {
direction: Direction::Forward,
contiguous_part_fn: |fields: ReversibleFields| {
let each_len = fields.map(|field| field.encoded_len(&renamed));
quote! {
|instance, tm| { 0 #(+ #each_len)* }
}
},
oneof_part_fn: |field: &Field| {
let encoded_len = field.encoded_len(&renamed);
quote! {
|instance, tm| { #encoded_len }
}
},
part_fn_ty: quote!(fn(&#part_fn_instance_ty, &mut #tag_measurer_ty) -> usize),
invoke_parts: quote! {
let mut total_len = 0usize;
for (_, len_func) in parts {
total_len += (len_func.unwrap())(#sort_group_self, tm)
}
total_len
},
},
),
});
quote! {
{
let tm = &mut #tag_measurer_ty::new();
0 #(+ #chunks)*
}
}
}
pub fn encode(&self, instance: &FieldTarget) -> TokenStream {
let crate_ = crate_name();
let sort_group_instance;
let part_fn_instance_ty;
if instance.has_instance() {
sort_group_instance = instance.clone();
part_fn_instance_ty = quote!(Self);
} else {
sort_group_instance = FieldTarget::RefsInstance(quote!(refs));
part_fn_instance_ty = quote!(__BilrostRefs);
}
let sort_group_self = sort_group_instance.self_expr();
let renamed = sort_group_instance.rename(quote!(instance));
let chunks = self.chunks.iter().map(|chunk| match chunk {
AlwaysOrdered(field) => field.encode(instance),
SortGroup(parts) => process_sort_groups(
parts,
&sort_group_instance,
SortGroupConfig {
direction: Direction::Forward,
contiguous_part_fn: |fields: ReversibleFields| {
let each_field = fields.map(|field| field.encode(&renamed));
quote! {
|instance, buf, tw| { #(#each_field)* }
}
},
oneof_part_fn: |field: &Field| {
let encode = field.encode(&renamed);
quote! {
|instance, buf, tw| { #encode }
}
},
part_fn_ty: quote!(
fn(&#part_fn_instance_ty, &mut __B, &mut #crate_::encoding::TagWriter)
),
invoke_parts: quote! {
for (_, encode_func) in parts {
(encode_func.unwrap())(#sort_group_self, buf, tw);
}
},
},
),
});
quote! {
{
let tw = &mut #crate_::encoding::TagWriter::new();
#(#chunks)*
}
}
}
pub fn prepend(&self, instance: &FieldTarget) -> TokenStream {
let crate_ = crate_name();
let sort_group_instance;
let part_fn_instance_ty;
if instance.has_instance() {
sort_group_instance = instance.clone();
part_fn_instance_ty = quote!(Self);
} else {
sort_group_instance = FieldTarget::RefsInstance(quote!(refs));
part_fn_instance_ty = quote!(__BilrostRefs);
}
let sort_group_self = sort_group_instance.self_expr();
let renamed = sort_group_instance.rename(quote!(instance));
let chunks = self.chunks.iter().rev().map(|chunk| match chunk {
AlwaysOrdered(field) => field.prepend(instance),
SortGroup(parts) => process_sort_groups(
parts,
&sort_group_instance,
SortGroupConfig {
direction: Direction::Reverse,
contiguous_part_fn: |fields: ReversibleFields| {
let each_field = fields.map(|field| field.prepend(&renamed));
quote! {
|instance, buf, tw| { #(#each_field)* }
}
},
oneof_part_fn: |field: &Field| {
let prepend = field.prepend(&renamed);
quote! {
|instance, buf, tw| { #prepend }
}
},
part_fn_ty: quote!(fn(&#part_fn_instance_ty, &mut __B, &mut #crate_::encoding::TagRevWriter)),
invoke_parts: quote! {
for (_, prepend_func) in parts {
(prepend_func.unwrap())(#sort_group_self, buf, tw);
}
},
},
),
});
quote! {
{
let tw = &mut #crate_::encoding::TagRevWriter::new();
#(#chunks)*
tw.finalize(buf);
}
}
}
}
#[derive(Clone)]
pub enum FieldTarget {
MessageInstance(TokenStream),
FreeVariantFields,
BoundVariantFields,
RefsInstance(TokenStream),
}
impl FieldTarget {
pub fn free_field_ident(field: &Field) -> Ident {
parse_str::<Ident>(&format!("field_{tag}", tag = field.first_tag()))
.expect("bound field name didn't parse as an ident")
}
pub fn has_instance(&self) -> bool {
match self {
FieldTarget::MessageInstance(..) | FieldTarget::RefsInstance(..) => true,
FieldTarget::FreeVariantFields | FieldTarget::BoundVariantFields => false,
}
}
pub fn self_expr(&self) -> Option<TokenStream> {
match self {
FieldTarget::MessageInstance(instance) => Some(instance.clone()),
FieldTarget::RefsInstance(instance) => Some(quote!(&#instance)),
FieldTarget::FreeVariantFields | FieldTarget::BoundVariantFields => None,
}
}
pub fn const_field_ref(&self, field: &Field) -> TokenStream {
match self {
FieldTarget::MessageInstance(instance) => {
let field_ident = field.ident();
quote!(&#instance.#field_ident)
}
FieldTarget::FreeVariantFields => {
let field_ident = Self::free_field_ident(field);
quote!(&#field_ident)
}
FieldTarget::BoundVariantFields => Self::free_field_ident(field).to_token_stream(),
FieldTarget::RefsInstance(instance) => {
let field_ident = Self::free_field_ident(field);
quote!(#instance.#field_ident)
}
}
}
pub fn mut_field_ref(&self, field: &Field) -> TokenStream {
match self {
FieldTarget::MessageInstance(instance) => {
let field_ident = field.ident();
quote!(&mut #instance.#field_ident)
}
FieldTarget::FreeVariantFields => {
let field_ident = Self::free_field_ident(field);
quote!(&mut #field_ident)
}
FieldTarget::BoundVariantFields => Self::free_field_ident(field).to_token_stream(),
FieldTarget::RefsInstance(instance) => {
let field_ident = Self::free_field_ident(field);
quote!(#instance.#field_ident)
}
}
}
pub fn rename(&self, new_instance_ident: TokenStream) -> Self {
match self {
FieldTarget::MessageInstance(_) => FieldTarget::MessageInstance(new_instance_ident),
FieldTarget::FreeVariantFields | FieldTarget::BoundVariantFields => {
panic!("free variant fields have no instance to rename")
}
FieldTarget::RefsInstance(_) => FieldTarget::RefsInstance(new_instance_ident),
}
}
pub fn prelude_for<'a>(
&self,
fields: impl IntoIterator<Item = &'a Field>,
) -> Option<TokenStream> {
let FieldTarget::RefsInstance(instance_ident) = self else {
return None;
};
let fields: Vec<_> = fields.into_iter().collect();
let field_idents: Vec<_> = fields
.iter()
.map(|field| Self::free_field_ident(field))
.collect();
let field_types: Vec<_> = fields.iter().map(|field| field.ty()).collect();
Some(quote! {
struct __BilrostRefs<'__r> {
#(#field_idents: &'__r #field_types,)*
}
let #instance_ident = &mut __BilrostRefs { #(#field_idents),* };
})
}
}