use crate::hir::*;
use anyhow::Result;
use quote::{quote, ToTokens};
use std::collections::{HashMap, HashSet};
pub(super) struct AdtPatternInfo {
pub abc_to_children: HashMap<String, Vec<String>>,
pub child_to_parent: HashMap<String, String>,
}
pub(super) fn convert_classes_to_rust(
classes: &[HirClass],
type_mapper: &crate::type_mapper::TypeMapper,
vararg_functions: &std::collections::HashSet<String>, ) -> Result<(Vec<proc_macro2::TokenStream>, HashMap<String, String>)> {
let adt_info = detect_adt_patterns(classes);
let mut class_items = Vec::new();
let mut processed_classes: HashSet<String> = HashSet::new();
for class in classes {
if processed_classes.contains(&class.name) {
continue;
}
if let Some(children) = adt_info.abc_to_children.get(&class.name) {
if !children.is_empty() && !class.type_params.is_empty() {
let tokens = generate_adt_enum(class, children, classes, type_mapper)?;
class_items.push(tokens);
for child_name in children {
processed_classes.insert(child_name.clone());
}
processed_classes.insert(class.name.clone());
continue;
}
}
let items =
crate::direct_rules::convert_class_to_struct(class, type_mapper, vararg_functions)?;
for item in items {
let tokens = item.to_token_stream();
class_items.push(tokens);
}
}
Ok((class_items, adt_info.child_to_parent))
}
pub(super) fn detect_adt_patterns(classes: &[HirClass]) -> AdtPatternInfo {
let mut abc_to_children: HashMap<String, Vec<String>> = HashMap::new();
let class_names: HashSet<&str> = classes.iter().map(|c| c.name.as_str()).collect();
for class in classes {
for base in &class.base_classes {
let base_name = base.split('[').next().unwrap_or(base);
if class_names.contains(base_name) {
abc_to_children
.entry(base_name.to_string())
.or_default()
.push(class.name.clone());
}
}
}
abc_to_children.retain(|parent_name, _| {
classes
.iter()
.find(|c| c.name == *parent_name)
.map(|c| {
!c.type_params.is_empty()
&& c.base_classes
.iter()
.any(|b| b.contains("ABC") || b.contains("Generic"))
})
.unwrap_or(false)
});
let mut child_to_parent = HashMap::new();
for (parent, children) in &abc_to_children {
for child in children {
child_to_parent.insert(child.clone(), parent.clone());
}
}
AdtPatternInfo {
abc_to_children,
child_to_parent,
}
}
pub(super) fn generate_adt_enum(
parent: &HirClass,
children: &[String],
all_classes: &[HirClass],
type_mapper: &crate::type_mapper::TypeMapper,
) -> Result<proc_macro2::TokenStream> {
let safe_name = crate::direct_rules::safe_class_name(&parent.name);
let enum_name = syn::Ident::new(&safe_name, proc_macro2::Span::call_site());
let type_params: Vec<syn::Ident> = parent
.type_params
.iter()
.map(|tp| syn::Ident::new(tp, proc_macro2::Span::call_site()))
.collect();
let generics = if type_params.is_empty() {
quote! {}
} else {
quote! { <#(#type_params: Clone),*> }
};
let generics_no_bounds = if type_params.is_empty() {
quote! {}
} else {
quote! { <#(#type_params),*> }
};
let mut variants = Vec::new();
for child_name in children {
let child = all_classes.iter().find(|c| &c.name == child_name);
if let Some(child_class) = child {
let safe_variant = crate::direct_rules::safe_class_name(&child_class.name);
let variant_name = syn::Ident::new(&safe_variant, proc_macro2::Span::call_site());
let field_types: Vec<proc_macro2::TokenStream> = child_class
.fields
.iter()
.filter(|f| !f.is_class_var && f.name != "_phantom")
.map(|f| {
let rust_type = type_mapper.map_type(&f.field_type);
crate::direct_rules::rust_type_to_syn_type(&rust_type)
.map(|t| quote! { #t })
.unwrap_or_else(|_| quote! { () })
})
.collect();
if field_types.len() == 1 {
let ft = &field_types[0];
variants.push(quote! { #variant_name(#ft) });
} else if field_types.is_empty() {
variants.push(quote! { #variant_name });
} else {
variants.push(quote! { #variant_name(#(#field_types),*) });
}
}
}
let methods = generate_adt_methods(parent, children, all_classes, type_mapper)?;
let result = quote! {
#[derive(Debug, Clone, PartialEq)]
pub enum #enum_name #generics {
#(#variants),*
}
impl #generics #enum_name #generics_no_bounds {
#methods
}
};
Ok(result)
}
pub(super) fn generate_adt_methods(
parent: &HirClass,
children: &[String],
all_classes: &[HirClass],
type_mapper: &crate::type_mapper::TypeMapper,
) -> Result<proc_macro2::TokenStream> {
let _type_params: Vec<syn::Ident> = parent
.type_params
.iter()
.map(|tp| syn::Ident::new(tp, proc_macro2::Span::call_site()))
.collect();
let mut methods = Vec::new();
for child_name in children {
let safe_variant = crate::direct_rules::safe_class_name(child_name);
let variant_name = syn::Ident::new(&safe_variant, proc_macro2::Span::call_site());
let method_name_str = format!("is_{}", safe_variant.to_lowercase());
let method_name = syn::Ident::new(&method_name_str, proc_macro2::Span::call_site());
methods.push(quote! {
pub fn #method_name(&self) -> bool {
matches!(self, Self::#variant_name(..))
}
});
}
for child_name in children {
let child = all_classes.iter().find(|c| &c.name == child_name);
if let Some(child_class) = child {
let safe_variant = crate::direct_rules::safe_class_name(&child_class.name);
let variant_name = syn::Ident::new(&safe_variant, proc_macro2::Span::call_site());
let method_name_str = format!("new_{}", safe_variant.to_lowercase());
let method_name = syn::Ident::new(&method_name_str, proc_macro2::Span::call_site());
let fields: Vec<_> = child_class
.fields
.iter()
.filter(|f| !f.is_class_var && f.name != "_phantom")
.collect();
if fields.len() == 1 {
let field = &fields[0];
let field_name = syn::Ident::new(&field.name, proc_macro2::Span::call_site());
let rust_type = type_mapper.map_type(&field.field_type);
let field_type = crate::direct_rules::rust_type_to_syn_type(&rust_type)?;
methods.push(quote! {
pub fn #method_name(#field_name: #field_type) -> Self {
Self::#variant_name(#field_name)
}
});
}
}
}
Ok(quote! { #(#methods)* })
}