use crate::{Ctx, ScalarBacking, backing_scalar};
use proc_macro2::TokenStream;
use quote::quote;
use ridl_ir::v2;
use std::collections::HashSet;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct Eligibility {
copy: bool,
eq: bool,
}
impl Eligibility {
const NONE: Self = Self {
copy: false,
eq: false,
};
const ALL: Self = Self {
copy: true,
eq: true,
};
const COPY_ONLY: Self = Self {
copy: true,
eq: false,
};
const EQ_ONLY: Self = Self {
copy: false,
eq: true,
};
fn meet(self, other: Self) -> Self {
Self {
copy: self.copy && other.copy,
eq: self.eq && other.eq,
}
}
}
pub(crate) fn derive_attr(ctx: &Ctx, decl: &v2::Decl) -> TokenStream {
let mut seen = HashSet::new();
let eligibility = decl_eligibility(ctx, decl, &mut seen);
attr(eligibility, numeric_named_scalar(decl))
}
pub(crate) fn tuple_derive_attr(ctx: &Ctx, tuple: &v2::TupleType) -> TokenStream {
let mut seen = HashSet::new();
attr(tuple_eligibility(ctx, tuple, &mut seen), false)
}
fn attr(eligibility: Eligibility, ordered: bool) -> TokenStream {
let mut traits = vec![quote! { Debug }, quote! { Clone }];
if eligibility.copy {
traits.push(quote! { Copy });
}
traits.push(quote! { PartialEq });
if eligibility.eq {
traits.push(quote! { Eq });
traits.push(quote! { Hash });
}
if ordered {
traits.push(quote! { PartialOrd });
if eligibility.eq {
traits.push(quote! { Ord });
}
}
quote! { #[derive(#(#traits),*)] }
}
fn numeric_named_scalar(decl: &v2::Decl) -> bool {
let Some(v2::decl::Kind::TypeDef(td)) = &decl.kind else {
return false;
};
matches!(
backing_scalar(td),
ScalarBacking::Float | ScalarBacking::Integer
)
}
fn decl_eligibility(ctx: &Ctx, decl: &v2::Decl, seen: &mut HashSet<String>) -> Eligibility {
if !seen.insert(decl.name.clone()) {
return Eligibility::NONE;
}
let result = match &decl.kind {
Some(v2::decl::Kind::TypeDef(td)) => scalar_eligibility(td),
Some(v2::decl::Kind::EnumDef(_)) | Some(v2::decl::Kind::EnumSetDef(_)) => Eligibility::ALL,
Some(v2::decl::Kind::StructDef(sd)) => sd
.members
.iter()
.filter_map(|member| match &member.member {
Some(v2::struct_member::Member::Field(field)) => Some(field),
Some(v2::struct_member::Member::Reserved(_)) | None => None,
})
.fold(Eligibility::ALL, |acc, field| {
acc.meet(field_eligibility(ctx, field, seen))
}),
Some(v2::decl::Kind::UnionDef(ud)) => ud.arms.iter().fold(Eligibility::ALL, |acc, arm| {
acc.meet(type_ref_eligibility(ctx, &arm.type_ref, seen))
}),
Some(_) | None => Eligibility::NONE,
};
seen.remove(&decl.name);
result
}
fn scalar_eligibility(td: &v2::TypeDef) -> Eligibility {
match backing_scalar(td) {
ScalarBacking::Float => Eligibility::COPY_ONLY,
ScalarBacking::Integer | ScalarBacking::Boolean => Eligibility::ALL,
ScalarBacking::String | ScalarBacking::Bytes => Eligibility::EQ_ONLY,
}
}
fn field_eligibility(ctx: &Ctx, field: &v2::Field, seen: &mut HashSet<String>) -> Eligibility {
match field.r#type.as_ref() {
Some(ft) => field_type_eligibility(ctx, ft, seen),
None => Eligibility::NONE,
}
}
fn field_type_eligibility(
ctx: &Ctx,
ft: &v2::FieldType,
seen: &mut HashSet<String>,
) -> Eligibility {
match &ft.kind {
Some(v2::field_type::Kind::Named(reference)) => type_ref_eligibility(ctx, reference, seen),
Some(v2::field_type::Kind::Primitive(prim)) => primitive_eligibility(*prim),
Some(v2::field_type::Kind::InlineScalar(td)) => scalar_eligibility(td),
Some(v2::field_type::Kind::Tuple(tuple)) => tuple_eligibility(ctx, tuple, seen),
Some(v2::field_type::Kind::Array(array)) => {
let inner = array
.element
.as_deref()
.map(|element| field_type_eligibility(ctx, element, seen))
.unwrap_or(Eligibility::NONE);
Eligibility {
copy: false,
eq: inner.eq,
}
}
Some(v2::field_type::Kind::Map(map)) => {
let key = map
.key
.as_deref()
.map(|key| field_type_eligibility(ctx, key, seen))
.unwrap_or(Eligibility::NONE);
let value = map
.value
.as_deref()
.map(|value| field_type_eligibility(ctx, value, seen))
.unwrap_or(Eligibility::NONE);
Eligibility {
copy: false,
eq: key.eq && value.eq,
}
}
Some(v2::field_type::Kind::Stream(_)) | None => Eligibility::NONE,
}
}
fn tuple_eligibility(ctx: &Ctx, tuple: &v2::TupleType, seen: &mut HashSet<String>) -> Eligibility {
tuple.fields.iter().fold(Eligibility::ALL, |acc, field| {
let inner = field
.r#type
.as_ref()
.map(|ft| field_type_eligibility(ctx, ft, seen))
.unwrap_or(Eligibility::NONE);
acc.meet(inner)
})
}
fn type_ref_eligibility(ctx: &Ctx, reference: &str, seen: &mut HashSet<String>) -> Eligibility {
match ctx.lookup(reference) {
Some(decl) => decl_eligibility(ctx, decl, seen),
None => Eligibility::NONE,
}
}
fn primitive_eligibility(prim: i32) -> Eligibility {
match v2::PrimitiveType::try_from(prim).unwrap_or(v2::PrimitiveType::Unspecified) {
v2::PrimitiveType::Float => Eligibility::COPY_ONLY,
v2::PrimitiveType::Integer | v2::PrimitiveType::Boolean => Eligibility::ALL,
v2::PrimitiveType::String | v2::PrimitiveType::Bytes => Eligibility::EQ_ONLY,
v2::PrimitiveType::Unspecified => Eligibility::NONE,
}
}