use std::collections::HashSet;
use prebindgen_flat::types_util::ident;
use proc_macro2::TokenStream;
use quote::quote;
use crate::{
declared_target::check_declared_target,
registry::{Registry, TypeKey},
};
mod error;
mod plan;
pub use self::{
error::{ExpandDeclError, ExpandError},
plan::{FoldArg, FoldBuild, FoldLeaf, FoldPlan, FoldShape, FoldVariant},
};
#[derive(Clone)]
pub enum Variant {
Ctor(syn::Ident),
Identity,
}
#[derive(Clone)]
pub struct ConstructorDecl {
pub target: TypeKey,
pub variants: Vec<Variant>,
pub default: bool,
}
#[derive(Clone)]
pub enum ExpandSel {
TopLevel,
Subset(Vec<Variant>),
}
#[derive(Clone)]
pub struct ExpandDecl {
pub func: syn::Ident,
pub param: syn::Ident,
pub declared_target: Option<TypeKey>,
pub sel: ExpandSel,
}
#[derive(Clone, Default)]
pub struct Expansions {
pub constructors: Vec<ConstructorDecl>,
pub expands: Vec<ExpandDecl>,
pub skip_construct: std::collections::HashSet<(syn::Ident, syn::Ident)>,
}
fn validate_declarations(exp: &Expansions) -> Result<(), ExpandError> {
let mut entries: Vec<ExpandDeclError> = Vec::new();
let mut ctor_targets: HashSet<String> = HashSet::new();
for c in &exp.constructors {
let target = c.target.as_str().to_string();
if c.variants.is_empty() {
entries.push(ExpandDeclError::EmptyConstructor {
target: target.clone(),
});
}
if !ctor_targets.insert(target.clone()) {
entries.push(ExpandDeclError::DuplicateConstructor { target });
}
}
let mut expand_keys: HashSet<(String, String)> = HashSet::new();
for ed in &exp.expands {
if let ExpandSel::Subset(v) = &ed.sel {
if v.is_empty() {
entries.push(ExpandDeclError::EmptySubset {
func: ed.func.clone(),
param: ed.param.clone(),
});
}
}
if !expand_keys.insert((ed.func.to_string(), ed.param.to_string())) {
entries.push(ExpandDeclError::DuplicateExpand {
func: ed.func.clone(),
param: ed.param.clone(),
});
}
}
if entries.is_empty() {
Ok(())
} else {
Err(ExpandError::InvalidDeclarations { entries })
}
}
pub(crate) fn apply<M>(
registry: &mut Registry<M>,
exp: &Expansions,
declared_fns: &std::collections::HashSet<syn::Ident>,
accessor_fns: &std::collections::HashSet<syn::Ident>,
method_receivers: &std::collections::HashMap<syn::Ident, TypeKey>,
) -> Result<(), ExpandError> {
validate_declarations(exp)?;
let mut done: HashSet<(String, String)> = HashSet::new();
let mut skip_construct = exp.skip_construct.clone();
for ed in &exp.expands {
if accessor_fns.contains(&ed.func) {
return Err(ExpandError::ConstructOnAccessor {
func: ed.func.clone(),
});
}
if let Some(declared) = &ed.declared_target {
let param_ty = param_reading(registry, &ed.func, &ed.param)?;
let bare = constructed_value(¶m_ty).key();
if bare != *declared {
return Err(ExpandError::ParamTypeMismatch {
func: ed.func.clone(),
param: ed.param.clone(),
declared: declared.as_str().to_string(),
actual: bare.as_str().to_string(),
});
}
}
if let ExpandSel::Subset(v) = &ed.sel {
if matches!(v.as_slice(), [Variant::Identity]) {
skip_construct.insert((ed.func.clone(), ed.param.clone()));
done.insert((ed.func.to_string(), ed.param.to_string()));
continue;
}
}
process_expand(registry, exp, ed)?;
done.insert((ed.func.to_string(), ed.param.to_string()));
}
for c in &exp.constructors {
if !c.default {
continue;
}
let ckey = c.target.clone();
for func in declared_fns {
if accessor_fns.contains(func) {
continue;
}
let Some(params) = registry.flat().function(&func).map(|f| f.params.clone()) else {
continue;
};
let receiver_key = method_receivers.get(func);
let mut receiver_skipped = false;
for (pname, pty) in params.iter().map(|p| (p.name.clone(), p.ty.clone())) {
let bare_key = constructed_value(&pty).key();
if !receiver_skipped && receiver_key == Some(&bare_key) {
receiver_skipped = true;
continue;
}
if bare_key != ckey {
continue;
}
if skip_construct.contains(&(func.clone(), pname.clone())) {
continue;
}
if !done.insert((func.to_string(), pname.to_string())) {
continue;
}
let ed = ExpandDecl {
func: func.clone(),
param: pname,
declared_target: None,
sel: ExpandSel::TopLevel,
};
process_expand(registry, exp, &ed)?;
}
}
}
Ok(())
}
fn param_reading<M>(
registry: &Registry<M>,
func: &syn::Ident,
param: &syn::Ident,
) -> Result<prebindgen_flat::flat::TypeRef, ExpandError> {
registry
.flat()
.function(&func)
.ok_or_else(|| ExpandError::UnknownFunction(func.clone()))?
.params
.iter()
.find(|p| &p.name == param)
.map(|p| p.ty.clone())
.ok_or_else(|| ExpandError::UnknownParam(func.clone(), param.clone()))
}
fn process_expand<M>(
registry: &mut Registry<M>,
exp: &Expansions,
ed: &ExpandDecl,
) -> Result<(), ExpandError> {
let param_ty = param_reading(registry, &ed.func, &ed.param)?;
let (optional, by_ref, target) = constructed_value_layers(¶m_ty);
let target_key = target.key();
let variants = resolve_constructor(exp, registry, &target_key, ed)?;
let mut visited: HashSet<TypeKey> = HashSet::new();
let plan = build_plan(
exp,
registry,
ed,
optional,
by_ref,
&target,
&variants,
&mut visited,
)?;
for leaf in &plan.leaves {
registry.require_input(&leaf.ty);
}
registry
.expansion_plans
.insert((ed.func.clone(), ed.param.clone()), plan);
Ok(())
}
fn resolve_constructor<M>(
exp: &Expansions,
_registry: &Registry<M>,
target_key: &TypeKey,
ed: &ExpandDecl,
) -> Result<Vec<Variant>, ExpandError> {
match &ed.sel {
ExpandSel::Subset(variants) => Ok(variants.clone()),
ExpandSel::TopLevel => exp
.constructors
.iter()
.find(|c| c.target == *target_key)
.map(|c| c.variants.clone())
.ok_or_else(|| ExpandError::NoConstructor {
func: ed.func.clone(),
param: ed.param.clone(),
target: target_key.to_string(),
}),
}
}
fn ctor_signature<M>(
registry: &Registry<M>,
func: &syn::Ident,
expected: &TypeKey,
) -> Result<CtorSig, ExpandError> {
let f = registry
.flat()
.function(&func)
.ok_or_else(|| ExpandError::UnknownConstructor(func.clone()))?;
let params: Vec<(syn::Ident, prebindgen_flat::flat::TypeRef)> = f
.params
.iter()
.map(|p| (p.name.clone(), p.ty.clone()))
.collect();
let (target, fallible) = match f.ret.fallible_parts() {
Some((ok, _)) => (ok.key(), true),
None => (f.ret.key(), false),
};
check_declared_target(func, &target, expected)?;
Ok(CtorSig { params, fallible })
}
struct CtorSig {
params: Vec<(syn::Ident, prebindgen_flat::flat::TypeRef)>,
fallible: bool,
}
#[allow(clippy::too_many_arguments)]
fn build_plan<M>(
exp: &Expansions,
registry: &Registry<M>,
ed: &ExpandDecl,
optional: bool,
by_ref: bool,
target: &prebindgen_flat::flat::TypeRef,
variants: &[Variant],
visited: &mut HashSet<TypeKey>,
) -> Result<FoldPlan, ExpandError> {
let param = &ed.param;
let mut leaves: Vec<FoldLeaf> = Vec::new();
if optional {
let [Variant::Ctor(func)] = variants else {
visited.insert(target.key());
let prefix = param.to_string();
let (selector, fold_variants) = build_core(
exp,
registry,
ed,
target,
variants,
by_ref,
&prefix,
&mut leaves,
visited,
)?;
visited.remove(&target.key());
return Ok(FoldPlan {
target: target.clone(),
by_ref,
shape: FoldShape::Optional((), Box::new(FoldShape::Base)),
leaves,
selector,
present: None,
variants: fold_variants,
});
};
let sig = ctor_signature(registry, func, &target.key())?;
if sig.params.len() == 1 {
let (_pn, pty) = &sig.params[0];
leaves.push(FoldLeaf {
name: param.clone(),
ty: pty.optional(),
});
return Ok(FoldPlan {
target: target.clone(),
by_ref,
shape: FoldShape::Optional((), Box::new(FoldShape::Base)),
leaves,
selector: None,
present: None,
variants: vec![FoldVariant {
ctor: Some(func.clone()),
fallible: sig.fallible,
clone: false,
inputs: vec![FoldArg::Leaf(0, false)],
}],
});
}
leaves.push(FoldLeaf {
name: ident(&format!("{}_present", param)),
ty: prebindgen_flat::flat::TypeRef::scalar(prebindgen_flat::flat::ScalarKind::Bool),
});
let prefix = param.to_string();
let mut inputs = Vec::new();
for (pname, pty) in &sig.params {
let name = ident(&format!("{}_{}", prefix, pname));
let arg = build_arg(
exp,
registry,
ed,
pty,
name,
false,
&mut leaves,
visited,
)?;
if matches!(arg, FoldArg::Build(_)) {
return Err(ExpandError::UnsupportedOptional {
func: ed.func.clone(),
param: ed.param.clone(),
reason: "nested-buildable constructor arguments cannot be optional",
});
}
inputs.push(arg);
}
return Ok(FoldPlan {
target: target.clone(),
by_ref,
shape: FoldShape::Optional((), Box::new(FoldShape::Base)),
leaves,
selector: None,
present: Some(0),
variants: vec![FoldVariant {
ctor: Some(func.clone()),
fallible: sig.fallible,
clone: false,
inputs,
}],
});
}
visited.insert(target.key());
let prefix = param.to_string();
let (selector, fold_variants) = build_core(
exp,
registry,
ed,
target,
variants,
by_ref,
&prefix,
&mut leaves,
visited,
)?;
visited.remove(&target.key());
Ok(FoldPlan {
target: target.clone(),
by_ref,
shape: FoldShape::Base,
leaves,
selector,
present: None,
variants: fold_variants,
})
}
#[allow(clippy::too_many_arguments)]
fn build_core<M>(
exp: &Expansions,
registry: &Registry<M>,
ed: &ExpandDecl,
target: &prebindgen_flat::flat::TypeRef,
variants: &[Variant],
by_ref: bool,
prefix: &str,
leaves: &mut Vec<FoldLeaf>,
visited: &mut HashSet<TypeKey>,
) -> Result<(Option<usize>, Vec<FoldVariant>), ExpandError> {
if let [Variant::Ctor(func)] = variants {
let sig = ctor_signature(registry, func, &target.key())?;
let np = sig.params.len();
let mut args = Vec::new();
for (pname, pty) in &sig.params {
let name = if np == 1 {
ident(prefix)
} else {
ident(&format!("{}_{}", prefix, pname))
};
args.push(build_arg(
exp, registry, ed, pty, name, false, leaves, visited,
)?);
}
Ok((
None,
vec![FoldVariant {
ctor: Some(func.clone()),
fallible: sig.fallible,
clone: false,
inputs: args,
}],
))
} else {
let sel_idx = leaves.len();
leaves.push(FoldLeaf {
name: ident(&format!("{}_sel", prefix)),
ty: prebindgen_flat::flat::TypeRef::scalar(prebindgen_flat::flat::ScalarKind::I32),
});
let mut fold_variants: Vec<FoldVariant> = Vec::new();
for (vi, v) in variants.iter().enumerate() {
match v {
Variant::Ctor(func) => {
let sig = ctor_signature(registry, func, &target.key())?;
let np = sig.params.len();
let mut args = Vec::new();
for (pi, (_pname, pty)) in sig.params.iter().enumerate() {
let name = if np == 1 {
ident(&format!("{}_{}", prefix, vi))
} else {
ident(&format!("{}_{}_{}", prefix, vi, pi))
};
args.push(build_arg(
exp, registry, ed, pty, name, true, leaves, visited,
)?);
}
fold_variants.push(FoldVariant {
ctor: Some(func.clone()),
fallible: sig.fallible,
clone: false,
inputs: args,
});
}
Variant::Identity => {
let idx = leaves.len();
let leaf_ty = if by_ref {
target.borrowed().optional()
} else {
target.optional()
};
leaves.push(FoldLeaf {
name: ident(&format!("{}_{}", prefix, vi)),
ty: leaf_ty,
});
fold_variants.push(FoldVariant {
ctor: None,
fallible: false,
clone: by_ref,
inputs: vec![FoldArg::Leaf(idx, false)],
});
}
}
}
Ok((Some(sel_idx), fold_variants))
}
}
#[allow(clippy::too_many_arguments)]
fn build_arg<M>(
exp: &Expansions,
registry: &Registry<M>,
ed: &ExpandDecl,
pty: &prebindgen_flat::flat::TypeRef,
name: syn::Ident,
dispatched: bool,
leaves: &mut Vec<FoldLeaf>,
visited: &mut HashSet<TypeKey>,
) -> Result<FoldArg, ExpandError> {
let (popt, pby_ref, bare) = constructed_value_layers(pty);
let key = bare.key();
let canon = exp
.constructors
.iter()
.find(|c| c.target == key && !c.variants.is_empty());
if let Some(c) = canon {
if dispatched {
return Err(ExpandError::UnsupportedRecursive {
func: ed.func.clone(),
reason: "recursive input under a selector-dispatched constructor variant",
});
}
if popt {
return Err(ExpandError::UnsupportedRecursive {
func: ed.func.clone(),
reason: "recursive input on an Option<…> parameter",
});
}
if !visited.insert(key.clone()) {
return Err(ExpandError::InputCycle {
ty: key.to_string(),
});
}
let variants = c.variants.clone();
let (selector, vars) = build_core(
exp,
registry,
ed,
&bare,
&variants,
pby_ref,
&name.to_string(),
leaves,
visited,
)?;
visited.remove(&key);
Ok(FoldArg::Build(Box::new(FoldBuild {
target: bare.clone(),
by_ref: pby_ref,
selector,
variants: vars,
})))
} else {
let idx = leaves.len();
let passthrough = dispatched && popt;
leaves.push(FoldLeaf {
name,
ty: if dispatched && !passthrough {
pty.optional()
} else {
pty.clone()
},
});
Ok(FoldArg::Leaf(idx, passthrough))
}
}
impl From<crate::declared_target::TargetMismatch> for ExpandError {
fn from(m: crate::declared_target::TargetMismatch) -> Self {
ExpandError::TargetMismatch {
ctor: m.func,
produces: m.actual,
expected: m.expected,
}
}
}
pub fn emit_fold(
plan: &FoldPlan,
leaf_locals: &[syn::Ident],
qualify: &dyn Fn(&syn::Ident) -> syn::Path,
) -> syn::Expr {
fold_shape(&plan.shape, plan, leaf_locals, None, qualify)
}
fn fold_shape(
shape: &FoldShape,
plan: &FoldPlan,
leaf_locals: &[syn::Ident],
bound: Option<&syn::Ident>,
qualify: &dyn Fn(&syn::Ident) -> syn::Path,
) -> syn::Expr {
match shape {
FoldShape::Base => emit_core_construct(plan, leaf_locals, bound, qualify),
FoldShape::Optional((), inner) => {
if let Some(sidx) = plan.selector {
let sel_local = &leaf_locals[sidx];
let inner_expr = emit_core_construct(plan, leaf_locals, None, qualify);
syn::parse_quote!(if #sel_local < 0 {
::core::result::Result::Ok(::core::option::Option::None)
} else {
(#inner_expr).map(::core::option::Option::Some)
})
} else if let Some(pidx) = plan.present {
let present_local = &leaf_locals[pidx];
let inner_expr = emit_core_construct(plan, leaf_locals, None, qualify);
syn::parse_quote!(if #present_local {
(#inner_expr).map(::core::option::Option::Some)
} else {
::core::result::Result::Ok(::core::option::Option::None)
})
} else {
let value = bound.unwrap_or(&leaf_locals[0]);
let inner_ident = ident("__inner");
let inner_expr = fold_shape(inner, plan, leaf_locals, Some(&inner_ident), qualify);
syn::parse_quote!(match #value {
::core::option::Option::Some(#inner_ident) => {
(#inner_expr).map(::core::option::Option::Some)
}
::core::option::Option::None => {
::core::result::Result::Ok(::core::option::Option::None)
}
})
}
}
FoldShape::Iterable(inner) => {
let value = bound.unwrap_or(&leaf_locals[0]);
let elem_ident = ident("__elem");
let inner_expr = fold_shape(inner, plan, leaf_locals, Some(&elem_ident), qualify);
syn::parse_quote!(
#value
.into_iter()
.map(|#elem_ident| #inner_expr)
.collect::<::core::result::Result<::std::vec::Vec<_>, _>>()
)
}
}
}
fn emit_core_construct(
plan: &FoldPlan,
leaf_locals: &[syn::Ident],
bound: Option<&syn::Ident>,
qualify: &dyn Fn(&syn::Ident) -> syn::Path,
) -> syn::Expr {
if let Some(v) = bound {
let var = &plan.variants[0];
let func = var
.ctor
.as_ref()
.expect("shaped expansion is single-constructor (never identity)");
return ctor_call_result(&qualify(func), std::slice::from_ref(v), var.fallible);
}
emit_dispatch(plan.selector, &plan.variants, leaf_locals, qualify)
}
fn emit_dispatch(
selector: Option<usize>,
variants: &[FoldVariant],
leaf_locals: &[syn::Ident],
qualify: &dyn Fn(&syn::Ident) -> syn::Path,
) -> syn::Expr {
match selector {
None => variant_result_expr(
&variants[0],
leaf_locals,
qualify,
false,
),
Some(si) => {
let sel = &leaf_locals[si];
let arms: Vec<TokenStream> = variants
.iter()
.enumerate()
.map(|(vi, v)| {
let lit = vi as i32;
let body =
variant_result_expr(v, leaf_locals, qualify, true);
quote!(#lit => #body,)
})
.collect();
syn::parse_quote!({
match #sel {
#(#arms)*
__sel => ::core::result::Result::Err(::std::format!(
"invalid constructor selector: {}",
__sel
)),
}
})
}
}
}
fn emit_build(
b: &FoldBuild,
leaf_locals: &[syn::Ident],
qualify: &dyn Fn(&syn::Ident) -> syn::Path,
) -> syn::Expr {
emit_dispatch(b.selector, &b.variants, leaf_locals, qualify)
}
fn variant_result_expr(
v: &FoldVariant,
leaf_locals: &[syn::Ident],
qualify: &dyn Fn(&syn::Ident) -> syn::Path,
dispatched: bool,
) -> syn::Expr {
let leaf = |a: &FoldArg| -> &syn::Ident {
match a {
FoldArg::Leaf(i, _) => &leaf_locals[*i],
FoldArg::Build(_) => {
unreachable!("recursive Build arg only in a non-dispatched single constructor")
}
}
};
match &v.ctor {
None => {
let loc = leaf(&v.inputs[0]);
let some_val: syn::Expr = if v.clone {
syn::parse_quote!(::core::result::Result::Ok(::core::clone::Clone::clone(
&*__v
)))
} else {
syn::parse_quote!(::core::result::Result::Ok(__v))
};
if dispatched {
syn::parse_quote!(match #loc {
::core::option::Option::Some(__v) => #some_val,
::core::option::Option::None => ::core::result::Result::Err(
::std::string::String::from("identity variant value missing")
),
})
} else if v.clone {
syn::parse_quote!(::core::result::Result::Ok(::core::clone::Clone::clone(&*#loc)))
} else {
syn::parse_quote!(::core::result::Result::Ok(#loc))
}
}
Some(func) => {
let path = qualify(func);
if dispatched {
let mut wrapped_locals: Vec<&syn::Ident> = Vec::new();
let mut wrapped_binds: Vec<syn::Ident> = Vec::new();
let mut call_args: Vec<syn::Expr> = Vec::new();
for (i, a) in v.inputs.iter().enumerate() {
let loc = leaf(a);
if matches!(a, FoldArg::Leaf(_, true)) {
call_args.push(syn::parse_quote!(#loc));
} else {
let b = ident(&format!("__p{}", i));
wrapped_locals.push(loc);
wrapped_binds.push(b.clone());
call_args.push(syn::parse_quote!(#b));
}
}
let call = ctor_call_result(&path, &call_args, v.fallible);
let missing = quote!(::core::result::Result::Err(::std::string::String::from(
"constructor variant input missing"
)));
match wrapped_locals.len() {
0 => call,
1 => {
let loc = wrapped_locals[0];
let p0 = &wrapped_binds[0];
syn::parse_quote!(match #loc {
::core::option::Option::Some(#p0) => #call,
::core::option::Option::None => #missing,
})
}
_ => {
let some_pats: Vec<TokenStream> = wrapped_binds
.iter()
.map(|b| quote!(::core::option::Option::Some(#b)))
.collect();
syn::parse_quote!(match ( #(#wrapped_locals),* ) {
( #(#some_pats),* ) => #call,
_ => #missing,
})
}
}
} else if v.inputs.iter().all(|a| matches!(a, FoldArg::Leaf(..))) {
let args: Vec<&syn::Ident> = v.inputs.iter().map(&leaf).collect();
ctor_call_result(&path, &args, v.fallible)
} else {
let mut stmts: Vec<TokenStream> = Vec::new();
let mut args: Vec<TokenStream> = Vec::new();
for (i, a) in v.inputs.iter().enumerate() {
let ai = ident(&format!("__a{}", i));
match a {
FoldArg::Leaf(li, _) => {
let loc = &leaf_locals[*li];
stmts.push(quote!(let #ai = #loc;));
args.push(quote!(#ai));
}
FoldArg::Build(b) => {
let be = emit_build(b, leaf_locals, qualify);
stmts.push(quote!(
let #ai = {
let __r: ::core::result::Result<_, ::std::string::String> = #be;
__r?
};
));
if b.by_ref {
args.push(quote!(&#ai));
} else {
args.push(quote!(#ai));
}
}
}
}
let call = ctor_call_result(&path, &args, v.fallible);
syn::parse_quote!({
(|| -> ::core::result::Result<_, ::std::string::String> {
#(#stmts)*
#call
})()
})
}
}
}
}
fn ctor_call_result<I: quote::ToTokens>(path: &syn::Path, args: &[I], fallible: bool) -> syn::Expr {
if fallible {
syn::parse_quote!(#path( #(#args),* ).map_err(|__e| ::std::format!("{}", __e)))
} else {
syn::parse_quote!(::core::result::Result::Ok(#path( #(#args),* )))
}
}
fn constructed_value(reading: &prebindgen_flat::flat::TypeRef) -> &prebindgen_flat::flat::TypeRef {
let after_opt = reading.optional_inner().unwrap_or(reading);
after_opt.borrow_target().unwrap_or(after_opt)
}
fn constructed_value_layers(
reading: &prebindgen_flat::flat::TypeRef,
) -> (bool, bool, prebindgen_flat::flat::TypeRef) {
let optional = reading.optional_inner().is_some();
let after_opt = reading.optional_inner().unwrap_or(reading);
let by_ref = after_opt.borrow_target().is_some();
let core = after_opt.borrow_target().unwrap_or(after_opt);
(optional, by_ref, core.clone())
}
#[cfg(test)]
mod tests;