use std::collections::{BTreeMap, BTreeSet};
use cratestack_core::{Field, TypeArity};
use quote::quote;
use crate::shared::ident;
use super::scalar::{domain_from_wire_expr, scalar_wire};
pub(super) struct RenderedUpdateMessage {
pub(super) tokens: proc_macro2::TokenStream,
}
pub(super) fn render_update_message(
message_name: &str,
domain_path: proc_macro2::TokenStream,
fields: &[&Field],
numbers: &BTreeMap<String, i32>,
enum_names: &BTreeSet<&str>,
) -> Result<RenderedUpdateMessage, String> {
let ident_tok = ident(message_name);
let mut prost_fields = Vec::with_capacity(fields.len());
let mut try_from_wire_lets = Vec::with_capacity(fields.len());
let mut try_from_wire_inits = Vec::with_capacity(fields.len());
for field in fields {
let number = *numbers
.get(&field.name)
.ok_or_else(|| format!("no `.pb.lock` entry for `{message_name}.{}`", field.name))?;
let plan = render_patch_field(message_name, field, number, enum_names);
prost_fields.push(plan.prost_field);
try_from_wire_lets.push(plan.try_from_wire_let);
let field_ident = ident(&field.name);
try_from_wire_inits.push(quote! { #field_ident, });
}
let tokens = quote! {
#[derive(Clone, PartialEq, ::cratestack::grpc::prost::Message)]
pub struct #ident_tok {
#(#prost_fields)*
}
impl ::core::convert::TryFrom<#ident_tok> for #domain_path {
type Error = ::cratestack::CoolError;
fn try_from(value: #ident_tok) -> ::core::result::Result<Self, Self::Error> {
#(#try_from_wire_lets)*
Ok(Self {
#(#try_from_wire_inits)*
})
}
}
};
Ok(RenderedUpdateMessage { tokens })
}
struct PatchFieldPlan {
prost_field: proc_macro2::TokenStream,
try_from_wire_let: proc_macro2::TokenStream,
}
fn render_patch_field(
owner: &str,
field: &Field,
number: i32,
enum_names: &BTreeSet<&str>,
) -> PatchFieldPlan {
let field_ident = ident(&field.name);
let field_name = field.name.as_str();
let type_name = field.ty.name.as_str();
let number_lit = proc_macro2::Literal::i32_unsuffixed(number);
let arity = field.ty.arity;
if let Some(wire) = scalar_wire(type_name) {
let rust_inner = &wire.rust_type;
let kind = &wire.prost_kind;
let to_domain = move |expr| domain_from_wire_expr(type_name, expr, owner, field_name);
render_patch_field_generic(
&field_ident,
number_lit,
arity,
quote! { #kind, optional },
quote! { #kind, repeated },
rust_inner.clone(),
to_domain,
)
} else if enum_names.contains(type_name) {
let enum_ident = ident(type_name);
let domain_enum_path = quote! { super::super::#enum_ident };
render_patch_field_generic(
&field_ident,
number_lit,
arity,
quote! { int32, optional },
quote! { int32, repeated },
quote! { i32 },
move |expr| {
quote! { <#domain_enum_path as ::core::convert::TryFrom<i32>>::try_from(#expr) }
},
)
} else {
let message_ident = ident(type_name);
let domain_message_path = quote! { super::super::#message_ident };
render_patch_field_generic(
&field_ident,
number_lit,
arity,
quote! { message, optional, boxed },
quote! { message, repeated },
quote! { Box<#message_ident> },
move |expr| quote! { #domain_message_path::try_from(*(#expr)) },
)
}
}
fn render_patch_field_generic(
field_ident: &syn::Ident,
number_lit: proc_macro2::Literal,
arity: TypeArity,
optional_attr: proc_macro2::TokenStream,
repeated_attr: proc_macro2::TokenStream,
rust_inner: proc_macro2::TokenStream,
domain_expr: impl Fn(proc_macro2::TokenStream) -> proc_macro2::TokenStream,
) -> PatchFieldPlan {
if arity == TypeArity::List {
let to_domain = domain_expr(quote! { raw });
return PatchFieldPlan {
prost_field: quote! {
#[prost(#repeated_attr, tag = #number_lit)]
pub #field_ident: Vec<#rust_inner>,
},
try_from_wire_let: quote! {
let #field_ident = if value.#field_ident.is_empty() {
None
} else {
Some(value.#field_ident
.into_iter()
.map(|raw| -> ::core::result::Result<_, ::cratestack::CoolError> { #to_domain })
.collect::<::core::result::Result<Vec<_>, ::cratestack::CoolError>>()?)
};
},
};
}
let to_domain = domain_expr(quote! { raw });
let prost_field = quote! {
#[prost(#optional_attr, tag = #number_lit)]
pub #field_ident: Option<#rust_inner>,
};
if arity == TypeArity::Optional {
PatchFieldPlan {
prost_field,
try_from_wire_let: quote! {
let #field_ident = match value.#field_ident {
None => None,
Some(raw) => Some(Some(#to_domain?)),
};
},
}
} else {
PatchFieldPlan {
prost_field,
try_from_wire_let: quote! {
let #field_ident = value.#field_ident
.map(|raw| -> ::core::result::Result<_, ::cratestack::CoolError> { #to_domain })
.transpose()?;
},
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn field(name: &str, ty: &str, arity: TypeArity) -> Field {
Field {
docs: vec![],
name: name.to_owned(),
name_span: cratestack_core::SourceSpan {
start: 0,
end: 0,
line: 0,
},
ty: cratestack_core::TypeRef {
name: ty.to_owned(),
name_span: cratestack_core::SourceSpan {
start: 0,
end: 0,
line: 0,
},
arity,
generic_args: vec![],
},
attributes: Vec::new(),
span: cratestack_core::SourceSpan {
start: 0,
end: 0,
line: 0,
},
}
}
#[test]
fn required_field_wraps_in_single_option() {
let name_field = field("name", "String", TypeArity::Required);
let fields = vec![&name_field];
let numbers = BTreeMap::from([("name".to_owned(), 1)]);
let rendered = render_update_message(
"UpdateWidgetInput",
quote! { super::super::UpdateWidgetInput },
&fields,
&numbers,
&BTreeSet::new(),
)
.expect("should render");
let rendered_str = rendered.tokens.to_string();
assert!(rendered_str.contains("pub name : Option < String >"));
assert!(!rendered_str.contains("Some (Some ("));
}
#[test]
fn optional_field_double_wraps_on_decode() {
let email_field = field("email", "String", TypeArity::Optional);
let fields = vec![&email_field];
let numbers = BTreeMap::from([("email".to_owned(), 1)]);
let rendered = render_update_message(
"UpdateWidgetInput",
quote! { super::super::UpdateWidgetInput },
&fields,
&numbers,
&BTreeSet::new(),
)
.expect("should render");
let rendered_str = rendered.tokens.to_string();
assert!(rendered_str.contains("Some (Some ("));
assert!(!rendered_str.contains("Some (None)"));
}
#[test]
fn missing_lock_entry_is_reported_not_panicked() {
let name_field = field("name", "String", TypeArity::Required);
let fields = vec![&name_field];
let result = render_update_message(
"UpdateWidgetInput",
quote! { super::super::UpdateWidgetInput },
&fields,
&BTreeMap::new(),
&BTreeSet::new(),
);
assert!(result.is_err());
}
}