use std::{
collections::BTreeMap,
path::{Path, PathBuf},
};
use proc_macro2::TokenStream;
use crate::{
destination::Destination,
prebindgen::Prebindgen,
registry::{Registry, TypeEntry, TypeKey},
};
#[derive(Debug)]
pub enum WriteError {
BadTokens {
phase: &'static str,
source: syn::Error,
},
}
impl std::fmt::Display for WriteError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
WriteError::BadTokens { phase, source } => {
write!(
f,
"generated tokens from {} did not parse: {}",
phase, source
)
}
}
}
}
impl std::error::Error for WriteError {}
pub fn write_rust<P: AsRef<Path>, E: Prebindgen>(
registry: &Registry<E::Metadata>,
ext: &E,
out_path: P,
) -> Result<PathBuf, WriteError> {
let emit = prebindgen_flat::Emit::new();
let mut items: Vec<syn::Item> = Vec::new();
items.extend(ext.prerequisites(registry, &emit));
for (_, item_fn) in collect_converter_items(registry) {
items.push(syn::Item::Fn(item_fn));
}
let declared = registry.declared();
let declared_fns = &declared.functions;
let declared_types = &declared.types;
let flat = registry.flat();
items.extend(parse_items_from_tokens(
"on_function",
sorted_by_name(flat.functions().map(|f| (&f.name, f)))
.into_iter()
.filter(|(ident, _)| declared_fns.contains(*ident))
.map(|(_, item)| ext.on_function(item, registry, &emit)),
)?);
items.extend(parse_items_from_tokens(
"on_struct",
sorted_by_name(flat.types().filter_map(|t| match t {
prebindgen_flat::flat::Type::Struct(s) => Some((&s.name, s)),
_ => None,
}))
.into_iter()
.filter(|(ident, _)| declared_types.contains_key(&TypeKey::from_ident(ident)))
.map(|(_, item)| ext.on_struct(item, registry, &emit)),
)?);
items.extend(parse_items_from_tokens(
"on_enum",
sorted_by_name(flat.types().filter_map(|t| match t {
prebindgen_flat::flat::Type::Variant(v) => Some((&v.name, t)),
prebindgen_flat::flat::Type::Enum(e) => Some((&e.name, t)),
_ => None,
}))
.into_iter()
.filter(|(ident, _)| declared_types.contains_key(&TypeKey::from_ident(ident)))
.map(|(_, t)| match t {
prebindgen_flat::flat::Type::Variant(v) => ext.on_variant(v, registry, &emit),
prebindgen_flat::flat::Type::Enum(e) => ext.on_enum(e, registry, &emit),
_ => unreachable!("filtered to the two enum shapes above"),
}),
)?);
let declared_consts = &declared.consts;
items.extend(parse_items_from_tokens(
"on_const",
sorted_by_name(flat.constants().map(|c| (&c.name, c)))
.into_iter()
.filter(|(ident, _)| {
declared_consts
.as_ref()
.is_none_or(|set| set.contains(*ident))
})
.map(|(_, item)| ext.on_const(item, registry, &emit)),
)?);
for guard in flat.guards() {
items.push(syn::Item::Const(emit.guard(guard)));
}
for item in &mut items {
ext.post_process_item(item, registry, &emit);
}
let dest: Destination = items.into_iter().collect();
Ok(dest.write(out_path))
}
fn collect_converter_items<M>(registry: &Registry<M>) -> Vec<(syn::Ident, syn::ItemFn)> {
let mut by_name: BTreeMap<String, (syn::Ident, syn::ItemFn)> = BTreeMap::new();
let mut collect = |entry: &TypeEntry<M>| {
let name = entry.function.sig.ident.clone();
by_name
.entry(name.to_string())
.or_insert_with(|| (name, entry.function.clone()));
for stage in &entry.pre_stages {
let sname = stage.function.sig.ident.clone();
by_name
.entry(sname.to_string())
.or_insert_with(|| (sname, stage.function.clone()));
}
};
walk_resolved(®istry.input_types, |_, entry| collect(entry));
walk_resolved(®istry.output_types, |_, entry| collect(entry));
by_name.into_values().collect()
}
fn walk_resolved<M, F: FnMut(&TypeKey, &TypeEntry<M>)>(
table: &std::collections::HashMap<TypeKey, crate::registry::TypeCell<M>>,
mut f: F,
) {
let mut keys: Vec<&TypeKey> = table.keys().collect();
keys.sort_by(|a, b| a.as_str().cmp(b.as_str()));
for key in keys {
if let Some(entry) = table.get(key).and_then(|c| c.entry.as_ref()) {
f(key, entry);
}
}
}
fn sorted_by_name<'a, T>(
items: impl Iterator<Item = (&'a syn::Ident, &'a T)>,
) -> Vec<(&'a syn::Ident, &'a T)>
where
T: 'a,
{
let mut items: Vec<(&syn::Ident, &T)> = items.collect();
items.sort_by_key(|(left, _)| left.to_string());
items
}
fn parse_items_from_tokens<I: IntoIterator<Item = TokenStream>>(
phase: &'static str,
iter: I,
) -> Result<Vec<syn::Item>, WriteError> {
let mut out = Vec::new();
for ts in iter {
if ts.is_empty() {
continue;
}
let file: syn::File =
syn::parse2(ts.clone()).map_err(|source| WriteError::BadTokens { phase, source })?;
out.extend(file.items);
}
Ok(out)
}
#[cfg(test)]
mod tests;