use proc_macro2::{Span, TokenStream};
use quote::{format_ident, quote};
use syn::Ident;
use crate::registry::{self, MemoParam};
pub fn expand() -> Result<TokenStream, syn::Error> {
registry::with(|reg| {
let mut output = TokenStream::new();
let memos: Vec<MemoInfo> = reg
.memos
.iter()
.map(|m| build_memo_info(m, reg))
.collect::<Result<Vec<_>, _>>()?;
let mut atom_memos: std::collections::HashMap<String, Vec<MemoInfo>> =
std::collections::HashMap::new();
for m in &memos {
atom_memos
.entry(m.primary_atom.clone())
.or_default()
.push(m.clone());
}
for memo in &memos {
output.extend(generate_input_struct(memo));
}
for atom in ®.atoms {
let atom_ident = Ident::new(&atom.name, Span::call_site());
let memos = atom_memos.get(&atom.name).cloned().unwrap_or_default();
let state_name = format_ident!("__Drv{}State", atom.name);
let state_fields: Vec<TokenStream> = memos
.iter()
.flat_map(|m| {
let in_field = format_ident!("{}_input", m.fn_name);
let out_field = format_ident!("{}_output", m.fn_name);
let input_struct = input_struct_ident(&m.fn_name);
let output_ty: syn::Type =
syn::parse_str(&m.output_ty_tokens).expect("output type should parse");
vec![
quote! { #in_field: Option<#input_struct> },
quote! { #out_field: Option<#output_ty> },
]
})
.collect();
output.extend(quote! {
#[doc(hidden)]
#[derive(Default)]
pub struct #state_name {
#(#state_fields,)*
}
impl ::drv::Atom for #atom_ident {
type State = #state_name;
}
});
}
for memo in &memos {
output.extend(generate_memo_fn(memo));
}
Ok(output)
})
}
fn input_struct_ident(fn_name: &str) -> Ident {
format_ident!("__Drv{}Input", snake_to_pascal(fn_name))
}
fn snake_to_pascal(s: &str) -> String {
s.split('_')
.map(|word| {
let mut chars = word.chars();
match chars.next() {
None => String::new(),
Some(c) => c.to_uppercase().collect::<String>() + chars.as_str(),
}
})
.collect()
}
fn generate_input_struct(memo: &MemoInfo) -> TokenStream {
let struct_ident = input_struct_ident(&memo.fn_name);
let fields: Vec<TokenStream> = memo
.params
.iter()
.map(|p| {
let name = Ident::new(&p.param_name, Span::call_site());
match &p.kind {
MemoParamKind::Lens { snapshot_ident, .. } => quote! { #name: #snapshot_ident },
MemoParamKind::Value { ty_tokens } => {
let ty: syn::Type = syn::parse_str(ty_tokens).expect("value type should parse");
quote! { #name: #ty }
}
MemoParamKind::ValueRef { referent_tokens } => {
let referent: syn::Type =
syn::parse_str(referent_tokens).expect("referent type should parse");
quote! { #name: <#referent as ::std::borrow::ToOwned>::Owned }
}
}
})
.collect();
quote! {
#[doc(hidden)]
#[derive(Default)]
pub struct #struct_ident {
#(pub #fields,)*
}
}
}
fn generate_memo_fn(memo: &MemoInfo) -> TokenStream {
let fn_ident = Ident::new(&memo.fn_name, Span::call_site());
let input_field = format_ident!("{}_input", memo.fn_name);
let output_field = format_ident!("{}_output", memo.fn_name);
let output_ty: syn::Type =
syn::parse_str(&memo.output_ty_tokens).expect("output type should parse");
let vis: TokenStream = memo.vis_tokens.parse().unwrap_or_else(|_| quote! {});
let body: TokenStream = memo.body_tokens.parse().expect("body should parse");
let input_struct = input_struct_ident(&memo.fn_name);
let compute_params: Vec<TokenStream> = memo
.params
.iter()
.map(|p| {
let pname = Ident::new(&p.param_name, Span::call_site());
match &p.kind {
MemoParamKind::Lens { lens_type, .. } => {
quote! { #pname: &#lens_type }
}
MemoParamKind::Value { ty_tokens } => {
let ty: syn::Type = syn::parse_str(ty_tokens).expect("value type should parse");
quote! { #pname: #ty }
}
MemoParamKind::ValueRef { referent_tokens } => {
let referent: syn::Type =
syn::parse_str(referent_tokens).expect("referent type should parse");
quote! { #pname: &#referent }
}
}
})
.collect();
let mut lifetime_params: Vec<TokenStream> = Vec::new();
let outer_params: Vec<TokenStream> = memo
.params
.iter()
.enumerate()
.map(|(i, p)| {
let pname = Ident::new(&p.param_name, Span::call_site());
match &p.kind {
MemoParamKind::Lens {
lens_name,
is_identity,
atom_name,
..
} => {
if *is_identity {
let atom_ident = Ident::new(atom_name, Span::call_site());
quote! { #pname: &#atom_ident }
} else {
let lens_ident = Ident::new(lens_name, Span::call_site());
let lt = syn::Lifetime::new(&format!("'drv{}", i), Span::call_site());
lifetime_params.push(quote! { #lt });
quote! { #pname: impl ::core::convert::Into<#lens_ident<#lt>> }
}
}
MemoParamKind::Value { ty_tokens } => {
let ty: syn::Type = syn::parse_str(ty_tokens).expect("value type should parse");
quote! { #pname: #ty }
}
MemoParamKind::ValueRef { referent_tokens } => {
let referent: syn::Type =
syn::parse_str(referent_tokens).expect("referent type should parse");
quote! { #pname: &#referent }
}
}
})
.collect();
let generics = if lifetime_params.is_empty() {
quote! {}
} else {
quote! { <#(#lifetime_params),*> }
};
let conversions: Vec<TokenStream> = memo
.params
.iter()
.map(|p| {
let pname = Ident::new(&p.param_name, Span::call_site());
match &p.kind {
MemoParamKind::Lens { is_identity, .. } if *is_identity => {
quote! { let #pname = #pname; }
}
MemoParamKind::Lens { .. } => {
quote! { let #pname: _ = #pname.into(); }
}
MemoParamKind::Value { .. } => {
quote! {}
}
MemoParamKind::ValueRef { .. } => {
quote! {}
}
}
})
.collect();
let first_lens = memo
.params
.iter()
.find(|p| matches!(p.kind, MemoParamKind::Lens { .. }))
.expect("memo must have at least one lens param (validated in memo.rs)");
let first_lens_name = Ident::new(&first_lens.param_name, Span::call_site());
let cache_expr = match &first_lens.kind {
MemoParamKind::Lens { is_identity, .. } if *is_identity => {
quote! { &#first_lens_name.__drv }
}
_ => quote! { #first_lens_name.__drv },
};
let fresh_checks: Vec<TokenStream> = memo
.params
.iter()
.map(|p| {
let pname = Ident::new(&p.param_name, Span::call_site());
let field = Ident::new(&p.param_name, Span::call_site());
match &p.kind {
MemoParamKind::Lens { is_identity, .. } if *is_identity => {
quote! { *#pname == __prev.#field }
}
MemoParamKind::Lens { .. } => {
quote! { #pname == __prev.#field }
}
MemoParamKind::Value { .. } => {
quote! { {
use ::drv::FastEqFallback as _;
::drv::FastEq(&#pname).fast_eq(&__prev.#field)
} }
}
MemoParamKind::ValueRef { .. } => {
quote! { #pname == __prev.#field }
}
}
})
.collect();
let snapshot_fields: Vec<TokenStream> = memo
.params
.iter()
.map(|p| {
let pname = Ident::new(&p.param_name, Span::call_site());
let field = Ident::new(&p.param_name, Span::call_site());
match &p.kind {
MemoParamKind::Lens {
is_identity,
atom_name,
fields: atom_fields,
..
} if *is_identity => {
let snap = format_ident!("__Drv{}Identity", atom_name);
let field_clones: Vec<TokenStream> = atom_fields
.iter()
.map(|f| {
let fname = Ident::new(&f.name, Span::call_site());
quote! { #fname: #pname.#fname.clone() }
})
.collect();
quote! { #field: #snap { #(#field_clones),* } }
}
MemoParamKind::Lens { .. } => {
quote! { #field: #pname.__drv_snapshot() }
}
MemoParamKind::Value { .. } => {
quote! { #field: #pname }
}
MemoParamKind::ValueRef { referent_tokens } => {
let referent: syn::Type =
syn::parse_str(referent_tokens).expect("referent type should parse");
quote! { #field: <#referent as ::std::borrow::ToOwned>::to_owned(#pname) }
}
}
})
.collect();
let compute_args: Vec<TokenStream> = memo
.params
.iter()
.map(|p| {
let pname = Ident::new(&p.param_name, Span::call_site());
match &p.kind {
MemoParamKind::Lens { is_identity, .. } if *is_identity => {
quote! { #pname }
}
MemoParamKind::Lens { .. } => {
quote! { &#pname }
}
MemoParamKind::Value { .. } => {
quote! { #pname.clone() }
}
MemoParamKind::ValueRef { .. } => {
quote! { #pname }
}
}
})
.collect();
quote! {
#vis fn #fn_ident #generics (#(#outer_params),*) -> #output_ty {
fn __compute(#(#compute_params),*) -> #output_ty {
#body
}
#(#conversions)*
{
let __state = ::core::cell::RefCell::borrow(&#cache_expr.inner);
if let Some(__prev) = __state.#input_field.as_ref() {
if #(#fresh_checks)&&* {
return <#output_ty as ::core::clone::Clone>::clone(
__state.#output_field.as_ref().unwrap()
);
}
}
}
let __out = __compute(#(#compute_args),*);
{
let mut __state = ::core::cell::RefCell::borrow_mut(&#cache_expr.inner);
__state.#output_field = Some(<#output_ty as ::core::clone::Clone>::clone(&__out));
__state.#input_field = Some(#input_struct {
#(#snapshot_fields,)*
});
}
__out
}
}
}
fn build_memo_info(
memo: ®istry::MemoRegistration,
reg: ®istry::Registry,
) -> Result<MemoInfo, syn::Error> {
let mut params = Vec::new();
let mut primary_atom: Option<String> = None;
for p in &memo.params {
match p {
MemoParam::Lens {
param_name,
lens_name,
} => {
let lens = reg
.find_lens(lens_name)
.expect("lens should exist (validated during #[drv::memo])");
if primary_atom.is_none() {
primary_atom = Some(lens.atom_name.clone());
}
let snapshot_ident = if lens.is_identity {
format_ident!("__Drv{}Identity", lens.atom_name)
} else {
format_ident!("__Drv{}", lens.name)
};
let lens_type = if lens.is_identity {
Ident::new(&lens.atom_name, Span::call_site()).into_token_stream()
} else {
let i = Ident::new(&lens.name, Span::call_site());
quote! { #i<'_> }
};
params.push(MemoParamLocal {
param_name: param_name.clone(),
kind: MemoParamKind::Lens {
lens_name: lens.name.clone(),
atom_name: lens.atom_name.clone(),
is_identity: lens.is_identity,
snapshot_ident,
lens_type,
fields: lens.fields.clone(),
},
});
}
MemoParam::Value {
param_name,
ty_tokens,
} => {
params.push(MemoParamLocal {
param_name: param_name.clone(),
kind: MemoParamKind::Value {
ty_tokens: ty_tokens.clone(),
},
});
}
MemoParam::ValueRef {
param_name,
referent_tokens,
} => {
params.push(MemoParamLocal {
param_name: param_name.clone(),
kind: MemoParamKind::ValueRef {
referent_tokens: referent_tokens.clone(),
},
});
}
}
}
Ok(MemoInfo {
fn_name: memo.fn_name.clone(),
vis_tokens: memo.vis_tokens.clone(),
output_ty_tokens: memo.output_ty_tokens.clone(),
body_tokens: memo.body_tokens.clone(),
params,
primary_atom: primary_atom.expect("memo must have at least one lens param"),
})
}
#[derive(Clone)]
struct MemoInfo {
fn_name: String,
vis_tokens: String,
output_ty_tokens: String,
body_tokens: String,
params: Vec<MemoParamLocal>,
primary_atom: String,
}
#[derive(Clone)]
struct MemoParamLocal {
param_name: String,
kind: MemoParamKind,
}
#[derive(Clone)]
enum MemoParamKind {
Lens {
lens_name: String,
atom_name: String,
is_identity: bool,
snapshot_ident: Ident,
lens_type: TokenStream,
fields: Vec<registry::LensField>,
},
Value {
ty_tokens: String,
},
ValueRef {
referent_tokens: String,
},
}
use quote::ToTokens;