use cairo_lang_syntax::node::db::SyntaxGroup;
use cairo_lang_syntax::node::helpers::{GenericParamEx, IsDependentType};
use cairo_lang_syntax::node::{Terminal, TypedSyntaxNode, ast};
use itertools::{Itertools, chain};
use smol_str::SmolStr;
pub struct MemberInfo {
pub name: SmolStr,
pub ty: String,
pub attributes: ast::AttributeList,
pub is_generics_dependent: bool,
}
impl MemberInfo {
pub fn impl_name(&self, trt: &str) -> String {
if self.is_generics_dependent {
let short_name = trt.split("::").last().unwrap_or(trt);
format!("__MEMBER_IMPL_{}_{short_name}", self.name)
} else {
format!("{}::<{}>", trt, self.ty)
}
}
pub fn drop_with(&self) -> String {
if self.is_generics_dependent {
format!("core::internal::DropWith::<{}, {}>", self.ty, self.impl_name("Drop"))
} else {
format!("core::internal::InferDrop::<{}>", self.ty)
}
}
pub fn destruct_with(&self) -> String {
if self.is_generics_dependent {
format!("core::internal::DestructWith::<{}, {}>", self.ty, self.impl_name("Destruct"))
} else {
format!("core::internal::InferDestruct::<{}>", self.ty)
}
}
}
pub enum TypeVariant {
Enum,
Struct,
}
pub struct GenericParamsInfo {
pub param_names: Vec<SmolStr>,
pub full_params: Vec<String>,
}
impl GenericParamsInfo {
pub fn new(db: &dyn SyntaxGroup, generic_params: ast::OptionWrappedGenericParamList) -> Self {
let ast::OptionWrappedGenericParamList::WrappedGenericParamList(gens) = generic_params
else {
return Self { param_names: Default::default(), full_params: Default::default() };
};
let params = gens.generic_params(db).elements(db);
Self {
param_names: params
.iter()
.map(|param| param.name(db).map(|n| n.text(db)).unwrap_or_else(|| "_".into()))
.collect(),
full_params: params
.iter()
.map(|param| param.as_syntax_node().get_text_without_trivia(db))
.collect(),
}
}
}
pub struct PluginTypeInfo {
pub name: SmolStr,
pub attributes: ast::AttributeList,
pub generics: GenericParamsInfo,
pub members_info: Vec<MemberInfo>,
pub type_variant: TypeVariant,
}
impl PluginTypeInfo {
pub fn new(db: &dyn SyntaxGroup, item_ast: &ast::ModuleItem) -> Option<Self> {
match item_ast {
ast::ModuleItem::Struct(struct_ast) => {
let generics = GenericParamsInfo::new(db, struct_ast.generic_params(db));
let members_info = extract_members(
db,
struct_ast.members(db),
&generics.param_names.iter().map(|p| p.as_str()).collect_vec(),
);
Some(Self {
name: struct_ast.name(db).text(db),
attributes: struct_ast.attributes(db),
generics,
members_info,
type_variant: TypeVariant::Struct,
})
}
ast::ModuleItem::Enum(enum_ast) => {
let generics = GenericParamsInfo::new(db, enum_ast.generic_params(db));
let members_info = extract_variants(
db,
enum_ast.variants(db),
&generics.param_names.iter().map(|p| p.as_str()).collect_vec(),
);
Some(Self {
name: enum_ast.name(db).text(db),
attributes: enum_ast.attributes(db),
generics,
members_info,
type_variant: TypeVariant::Enum,
})
}
_ => None,
}
}
pub fn impl_header(&self, derived_trait: &str, dependent_traits: &[&str]) -> String {
let derived_trait_name = derived_trait.split("::").last().unwrap_or(derived_trait);
format!(
"impl {name}{derived_trait_name}<{generics}> of {derived_trait}::<{full_typename}>",
name = self.name,
generics =
self.impl_generics(dependent_traits, |trt, ty| format!("{trt}<{ty}>")).join(", "),
full_typename = self.full_typename(),
)
}
pub fn impl_generics(
&self,
dependent_traits: &[&str],
dep_req: fn(&str, &str) -> String,
) -> Vec<String> {
chain!(
self.generics.full_params.iter().cloned(),
self.members_info.iter().filter(|m| m.is_generics_dependent).flat_map(|m| {
dependent_traits
.iter()
.cloned()
.map(move |trt| format!("impl {}: {}", m.impl_name(trt), dep_req(trt, &m.ty)))
})
)
.collect()
}
pub fn full_typename(&self) -> String {
if self.generics.param_names.is_empty() {
self.name.to_string()
} else {
format!("{}<{}>", self.name, self.generics.param_names.iter().join(", "))
}
}
}
fn extract_members(
db: &dyn SyntaxGroup,
members: ast::MemberList,
generics: &[&str],
) -> Vec<MemberInfo> {
members
.elements(db)
.into_iter()
.map(|member| MemberInfo {
name: member.name(db).text(db),
ty: member.type_clause(db).ty(db).as_syntax_node().get_text_without_trivia(db),
attributes: member.attributes(db),
is_generics_dependent: member.type_clause(db).ty(db).is_dependent_type(db, generics),
})
.collect()
}
fn extract_variants(
db: &dyn SyntaxGroup,
variants: ast::VariantList,
generics: &[&str],
) -> Vec<MemberInfo> {
variants
.elements(db)
.into_iter()
.map(|variant| MemberInfo {
name: variant.name(db).text(db),
ty: match variant.type_clause(db) {
ast::OptionTypeClause::Empty(_) => "()".to_string(),
ast::OptionTypeClause::TypeClause(t) => {
t.ty(db).as_syntax_node().get_text_without_trivia(db)
}
},
attributes: variant.attributes(db),
is_generics_dependent: match variant.type_clause(db) {
ast::OptionTypeClause::Empty(_) => false,
ast::OptionTypeClause::TypeClause(t) => t.ty(db).is_dependent_type(db, generics),
},
})
.collect()
}