use super::helpers::{extract_cfg_condition, is_test_gated};
use crate::core::ir::{DefaultValue, FieldDef, TypeRef};
use ahash::AHashMap;
use quote::ToTokens;
use syn;
mod function_default;
mod mutation;
mod struct_literal;
pub(crate) use function_default::{FreeFunctionIndex, collect_free_functions, fold_constant_default_functions};
use struct_literal::struct_expr_defaults;
pub(crate) type ConstructorIndex<'a> = AHashMap<(String, String), &'a syn::ImplItemFn>;
const MAX_DELEGATION_DEPTH: usize = 4;
pub(crate) fn extract_default_values(
item: &syn::ItemImpl,
self_type: &str,
fields: &mut [FieldDef],
literal_consts: &AHashMap<String, DefaultValue>,
constructors: &ConstructorIndex<'_>,
binding_excluded: bool,
) {
let default_fn = item.items.iter().find_map(|impl_item| {
if let syn::ImplItem::Fn(method) = impl_item
&& method.sig.ident == "default"
{
return Some(method);
}
None
});
let Some(default_fn) = default_fn else {
mark_unresolved(fields, "impl Default block without a `fn default()` item");
return;
};
let field_types: AHashMap<String, TypeRef> = fields
.iter()
.map(|field| (field.name.clone(), field.ty.clone()))
.collect();
let scope = EvalScope::new(self_type, literal_consts, &field_types);
let defaults = if let Some(body) = mutation::read_struct_body(&default_fn.block) {
mutation::struct_body_defaults(&body, &scope)
} else if let Some(delegated) = follow_delegation(&default_fn.block, self_type, constructors, &scope, 0) {
delegated
} else if tail_is_bare_self(&default_fn.block, self_type) {
AHashMap::new()
} else if let Some(single_field) = single_field_const_tail_default(&default_fn.block, fields, &scope) {
single_field
} else {
let body = default_fn.block.to_token_stream().to_string();
if !binding_excluded {
tracing::warn!(
target: "alef::extract::defaults",
rust_type = self_type,
body = %body,
"`impl Default` body is neither a struct literal nor a constant-foldable delegation; \
field defaults are unresolved"
);
}
mark_unresolved(fields, &body);
return;
};
for field in fields.iter_mut() {
if let Some(default_val) = defaults.get(&field.name) {
field.typed_default = Some(default_val.clone());
} else {
field.typed_default = Some(DefaultValue::Empty);
}
}
}
fn mark_unresolved(fields: &mut [FieldDef], body: &str) {
for field in fields.iter_mut() {
field.typed_default = Some(DefaultValue::Unresolved(body.to_string()));
}
}
pub(crate) fn collect_constructors(items: &[syn::Item]) -> ConstructorIndex<'_> {
let mut index = ConstructorIndex::new();
for item in items {
let syn::Item::Impl(item_impl) = item else {
continue;
};
if item_impl.trait_.is_some() || is_test_gated(&item_impl.attrs) {
continue;
}
let Some(type_name) = path_type_name(&item_impl.self_ty) else {
continue;
};
for impl_item in &item_impl.items {
let syn::ImplItem::Fn(method) = impl_item else {
continue;
};
if matches!(method.sig.inputs.first(), Some(syn::FnArg::Receiver(_))) {
continue;
}
index.insert((type_name.clone(), method.sig.ident.to_string()), method);
}
}
index
}
fn path_type_name(ty: &syn::Type) -> Option<String> {
match ty {
syn::Type::Path(path) => path.path.segments.last().map(|segment| segment.ident.to_string()),
_ => None,
}
}
pub(crate) fn collect_literal_consts(items: &[syn::Item]) -> AHashMap<String, DefaultValue> {
let mut consts = AHashMap::new();
for item in items {
match item {
syn::Item::Const(item_const) => {
if let Some(value) = const_literal_value(&item_const.expr) {
consts.insert(item_const.ident.to_string(), value);
}
}
syn::Item::Impl(item_impl) if item_impl.trait_.is_none() && !is_test_gated(&item_impl.attrs) => {
let Some(type_name) = path_type_name(&item_impl.self_ty) else {
continue;
};
for impl_item in &item_impl.items {
if let syn::ImplItem::Const(assoc_const) = impl_item
&& let Some(value) = const_literal_value(&assoc_const.expr)
{
consts.insert(format!("{type_name}::{}", assoc_const.ident), value);
}
}
}
_ => {}
}
}
consts
}
fn const_literal_value(expr: &syn::Expr) -> Option<DefaultValue> {
match expr {
syn::Expr::Lit(lit) => match &lit.lit {
syn::Lit::Str(s) => Some(DefaultValue::StringLiteral(s.value())),
syn::Lit::Char(c) => Some(DefaultValue::StringLiteral(c.value().to_string())),
syn::Lit::Bool(b) => Some(DefaultValue::BoolLiteral(b.value)),
syn::Lit::Int(i) => i.base10_parse::<i64>().ok().map(DefaultValue::IntLiteral),
syn::Lit::Float(f) => f.base10_parse::<f64>().ok().map(DefaultValue::FloatLiteral),
_ => None,
},
syn::Expr::Unary(unary) if matches!(unary.op, syn::UnOp::Neg(_)) => match const_literal_value(&unary.expr)? {
DefaultValue::IntLiteral(v) => Some(DefaultValue::IntLiteral(-v)),
DefaultValue::FloatLiteral(v) => Some(DefaultValue::FloatLiteral(-v)),
_ => None,
},
syn::Expr::Call(call) if call.args.len() == 1 => {
let syn::Expr::Path(path) = &*call.func else {
return None;
};
let name = path.path.segments.last()?.ident.to_string();
if !name.starts_with(|c: char| c.is_ascii_uppercase()) {
return None;
}
const_literal_value(call.args.first()?)
}
_ => None,
}
}
struct EvalScope<'a> {
self_type: &'a str,
literal_consts: &'a AHashMap<String, DefaultValue>,
field_types: &'a AHashMap<String, TypeRef>,
params: AHashMap<String, DefaultValue>,
}
impl<'a> EvalScope<'a> {
fn new(
self_type: &'a str,
literal_consts: &'a AHashMap<String, DefaultValue>,
field_types: &'a AHashMap<String, TypeRef>,
) -> Self {
Self {
self_type,
literal_consts,
field_types,
params: AHashMap::new(),
}
}
fn with_params(&self, params: AHashMap<String, DefaultValue>) -> EvalScope<'a> {
EvalScope {
self_type: self.self_type,
literal_consts: self.literal_consts,
field_types: self.field_types,
params,
}
}
fn associated_const(&self, owner: &str, name: &str) -> Option<DefaultValue> {
let owner = if owner == "Self" { self.self_type } else { owner };
self.literal_consts.get(&format!("{owner}::{name}")).cloned()
}
}
fn carries_value(value: &DefaultValue) -> bool {
matches!(
value,
DefaultValue::BoolLiteral(_)
| DefaultValue::StringLiteral(_)
| DefaultValue::IntLiteral(_)
| DefaultValue::FloatLiteral(_)
| DefaultValue::EnumVariant(_)
| DefaultValue::TupleVariant(_, _)
| DefaultValue::StructVariant(_, _)
| DefaultValue::ListLiteral(_)
)
}
fn follow_delegation(
block: &syn::Block,
self_type: &str,
constructors: &ConstructorIndex<'_>,
scope: &EvalScope<'_>,
depth: usize,
) -> Option<AHashMap<String, DefaultValue>> {
if depth >= MAX_DELEGATION_DEPTH {
return None;
}
let call = tail_call_expr(block)?;
let syn::Expr::Path(path) = &*call.func else {
return None;
};
let segments: Vec<String> = path.path.segments.iter().map(|s| s.ident.to_string()).collect();
let [owner, fn_name] = segments.as_slice() else {
return None;
};
if owner.as_str() != "Self" && owner.as_str() != self_type {
return None;
}
if fn_name.as_str() == "default" {
return None;
}
let target = constructors.get(&(self_type.to_string(), fn_name.clone()))?;
let mut params = AHashMap::new();
let mut arguments = call.args.iter();
for input in &target.sig.inputs {
let syn::FnArg::Typed(pat_type) = input else {
return None;
};
let argument = arguments.next()?;
let syn::Pat::Ident(pat_ident) = pat_type.pat.as_ref() else {
continue;
};
let value = expr_to_default_value(argument, scope, None);
if carries_value(&value) {
params.insert(pat_ident.ident.to_string(), value);
}
}
if arguments.next().is_some() {
return None;
}
let inner = scope.with_params(params);
if let Some(body) = mutation::read_struct_body(&target.block) {
return Some(mutation::struct_body_defaults(&body, &inner));
}
if tail_is_bare_self(&target.block, self_type) {
return Some(AHashMap::new());
}
follow_delegation(&target.block, self_type, constructors, &inner, depth + 1)
}
fn tail_call_expr(block: &syn::Block) -> Option<&syn::ExprCall> {
match block.stmts.last()? {
syn::Stmt::Expr(expr, _) => unwrap_to_call_expr(expr),
_ => None,
}
}
fn unwrap_to_call_expr(expr: &syn::Expr) -> Option<&syn::ExprCall> {
match expr {
syn::Expr::Call(call) => Some(call),
syn::Expr::Block(b) => tail_call_expr(&b.block),
syn::Expr::Return(ret) => ret.expr.as_deref().and_then(unwrap_to_call_expr),
_ => None,
}
}
fn single_field_const_tail_default(
block: &syn::Block,
fields: &[FieldDef],
scope: &EvalScope<'_>,
) -> Option<AHashMap<String, DefaultValue>> {
let [field] = fields else {
return None;
};
let path = tail_path_expr(block)?;
let segments: Vec<String> = path.path.segments.iter().map(|s| s.ident.to_string()).collect();
let [owner, name] = segments.as_slice() else {
return None;
};
let value = scope.associated_const(owner, name)?;
let mut defaults = AHashMap::new();
defaults.insert(field.name.clone(), value);
Some(defaults)
}
fn tail_path_expr(block: &syn::Block) -> Option<&syn::ExprPath> {
match block.stmts.last()? {
syn::Stmt::Expr(expr, _) => unwrap_to_path_expr(expr),
_ => None,
}
}
fn unwrap_to_path_expr(expr: &syn::Expr) -> Option<&syn::ExprPath> {
match expr {
syn::Expr::Path(path) => Some(path),
syn::Expr::Block(b) => tail_path_expr(&b.block),
syn::Expr::Return(ret) => ret.expr.as_deref().and_then(unwrap_to_path_expr),
_ => None,
}
}
fn tail_is_bare_self(block: &syn::Block, self_type: &str) -> bool {
let Some(path) = tail_path_expr(block) else {
return false;
};
if path.path.segments.len() != 1 {
return false;
}
let ident = &path.path.segments[0].ident;
ident == "Self" || ident == self_type
}
trait FieldMemberExt {
fn member_named(&self) -> Option<&syn::Ident>;
}
impl FieldMemberExt for syn::FieldValue {
fn member_named(&self) -> Option<&syn::Ident> {
match &self.member {
syn::Member::Named(ident) => Some(ident),
syn::Member::Unnamed(_) => None,
}
}
}
fn unreadable(expr: &syn::Expr) -> DefaultValue {
DefaultValue::Unresolved(expr.to_token_stream().to_string())
}
fn admits_enum_variant(field_ty: Option<&TypeRef>) -> bool {
match field_ty {
None | Some(TypeRef::Named(_)) => true,
Some(TypeRef::Optional(inner) | TypeRef::Vec(inner)) => admits_enum_variant(Some(&**inner)),
Some(_) => false,
}
}
fn expr_to_default_value(expr: &syn::Expr, scope: &EvalScope<'_>, field_ty: Option<&TypeRef>) -> DefaultValue {
match expr {
syn::Expr::Lit(lit) => match &lit.lit {
syn::Lit::Bool(b) => DefaultValue::BoolLiteral(b.value),
syn::Lit::Int(i) => {
if let Ok(val) = i.base10_parse::<i64>() {
DefaultValue::IntLiteral(val)
} else {
unreadable(expr)
}
}
syn::Lit::Float(f) => {
if let Ok(val) = f.base10_parse::<f64>() {
DefaultValue::FloatLiteral(val)
} else {
unreadable(expr)
}
}
syn::Lit::Char(c) => DefaultValue::StringLiteral(c.value().to_string()),
syn::Lit::Str(s) => DefaultValue::StringLiteral(s.value()),
_ => unreadable(expr),
},
syn::Expr::Reference(syn::ExprReference { expr: inner, .. })
| syn::Expr::Paren(syn::ExprParen { expr: inner, .. })
| syn::Expr::Group(syn::ExprGroup { expr: inner, .. }) => expr_to_default_value(inner, scope, field_ty),
syn::Expr::Unary(unary) if matches!(unary.op, syn::UnOp::Neg(_)) => {
match expr_to_default_value(&unary.expr, scope, field_ty) {
DefaultValue::IntLiteral(v) => DefaultValue::IntLiteral(-v),
DefaultValue::FloatLiteral(v) => DefaultValue::FloatLiteral(-v),
_ => unreadable(expr),
}
}
syn::Expr::Binary(bin) => {
let lhs = expr_to_default_value(&bin.left, scope, field_ty);
let rhs = expr_to_default_value(&bin.right, scope, field_ty);
match (lhs, rhs) {
(DefaultValue::IntLiteral(a), DefaultValue::IntLiteral(b)) => match bin.op {
syn::BinOp::Add(_) => a
.checked_add(b)
.map(DefaultValue::IntLiteral)
.unwrap_or_else(|| unreadable(expr)),
syn::BinOp::Sub(_) => a
.checked_sub(b)
.map(DefaultValue::IntLiteral)
.unwrap_or_else(|| unreadable(expr)),
syn::BinOp::Mul(_) => a
.checked_mul(b)
.map(DefaultValue::IntLiteral)
.unwrap_or_else(|| unreadable(expr)),
syn::BinOp::Div(_) if b != 0 => DefaultValue::IntLiteral(a / b),
syn::BinOp::Rem(_) if b != 0 => DefaultValue::IntLiteral(a % b),
syn::BinOp::Shl(_) if (0..63).contains(&b) => a
.checked_shl(b as u32)
.map(DefaultValue::IntLiteral)
.unwrap_or_else(|| unreadable(expr)),
syn::BinOp::Shr(_) if (0..63).contains(&b) => DefaultValue::IntLiteral(a >> (b as u32)),
syn::BinOp::BitOr(_) => DefaultValue::IntLiteral(a | b),
syn::BinOp::BitAnd(_) => DefaultValue::IntLiteral(a & b),
syn::BinOp::BitXor(_) => DefaultValue::IntLiteral(a ^ b),
_ => unreadable(expr),
},
(DefaultValue::FloatLiteral(a), DefaultValue::FloatLiteral(b)) => match bin.op {
syn::BinOp::Add(_) => DefaultValue::FloatLiteral(a + b),
syn::BinOp::Sub(_) => DefaultValue::FloatLiteral(a - b),
syn::BinOp::Mul(_) => DefaultValue::FloatLiteral(a * b),
syn::BinOp::Div(_) if b != 0.0 => DefaultValue::FloatLiteral(a / b),
_ => unreadable(expr),
},
_ => unreadable(expr),
}
}
syn::Expr::MethodCall(mc) => {
let method_name = mc.method.to_string();
match method_name.as_str() {
"to_string" | "to_owned" | "into" => {
if let syn::Expr::Lit(lit) = &*mc.receiver
&& let syn::Lit::Str(s) = &lit.lit
{
return DefaultValue::StringLiteral(s.value());
}
match resolve_ident(&mc.receiver, scope) {
Some(value @ DefaultValue::StringLiteral(_)) => value,
Some(value) if method_name == "into" => value,
_ => unreadable(expr),
}
}
_ => unreadable(expr),
}
}
syn::Expr::Call(call) => {
if let syn::Expr::Path(path) = &*call.func {
let segments: Vec<String> = path.path.segments.iter().map(|s| s.ident.to_string()).collect();
if (segments == ["Some"] || segments == ["Option", "Some"])
&& call.args.len() == 1
&& let Some(inner) = call.args.first()
{
return expr_to_default_value(inner, scope, field_ty);
}
if segments == ["String", "from"] && call.args.len() == 1 {
if let Some(syn::Expr::Lit(lit)) = call.args.first()
&& let syn::Lit::Str(s) = &lit.lit
{
return DefaultValue::StringLiteral(s.value());
}
if let Some(argument) = call.args.first()
&& let Some(value @ DefaultValue::StringLiteral(_)) = resolve_ident(argument, scope)
{
return value;
}
return unreadable(expr);
}
if segments == ["String", "new"] && call.args.is_empty() {
return DefaultValue::StringLiteral(String::new());
}
if let [.., owner, variant] = segments.as_slice()
&& owner == "Cow"
&& matches!(variant.as_str(), "Borrowed" | "Owned")
&& call.args.len() == 1
&& let Some(inner) = call.args.first()
{
return match expr_to_default_value(inner, scope, field_ty) {
DefaultValue::Unresolved(_) | DefaultValue::FunctionCall(_) => unreadable(expr),
resolved => resolved,
};
}
if segments.len() == 2 && segments[1] == "new" && call.args.is_empty() {
let type_name = &segments[0];
if matches!(
type_name.as_str(),
"Vec" | "HashMap" | "HashSet" | "BTreeMap" | "BTreeSet" | "AHashMap" | "AHashSet"
) {
return DefaultValue::Empty;
}
}
if segments == ["Duration", "from_secs"] && call.args.len() == 1 {
if let Some(syn::Expr::Lit(lit)) = call.args.first()
&& let syn::Lit::Int(i) = &lit.lit
&& let Ok(val) = i.base10_parse::<i64>()
{
return DefaultValue::IntLiteral(val * 1000);
}
return unreadable(expr);
}
if segments == ["Duration", "from_millis"] && call.args.len() == 1 {
if let Some(syn::Expr::Lit(lit)) = call.args.first()
&& let syn::Lit::Int(i) = &lit.lit
&& let Ok(val) = i.base10_parse::<i64>()
{
return DefaultValue::IntLiteral(val);
}
return unreadable(expr);
}
if segments.last().is_some_and(|s| s == "default") {
return DefaultValue::Empty;
}
if !call.args.is_empty()
&& let Some(variant) = segments.last()
&& variant.starts_with(|c: char| c.is_ascii_uppercase())
{
let mut values = Vec::with_capacity(call.args.len());
for argument in &call.args {
let value = expr_to_default_value(argument, scope, None);
if !carries_value(&value) {
return unreadable(expr);
}
values.push(value);
}
return DefaultValue::TupleVariant(variant.clone(), values);
}
if call.args.is_empty() {
return DefaultValue::FunctionCall(segments.join("::"));
}
}
unreadable(expr)
}
syn::Expr::Struct(struct_expr) => {
if struct_expr.rest.is_some() {
return unreadable(expr);
}
let Some(variant) = struct_expr.path.segments.last().map(|s| s.ident.to_string()) else {
return unreadable(expr);
};
let mut fields = Vec::with_capacity(struct_expr.fields.len());
for field_value in &struct_expr.fields {
let Some(name) = field_value.member_named() else {
return unreadable(expr);
};
let value = expr_to_default_value(&field_value.expr, scope, None);
if !carries_value(&value) {
return unreadable(expr);
}
fields.push((name.to_string(), value));
}
DefaultValue::StructVariant(variant, fields)
}
syn::Expr::Path(path) => {
if let Some(value) = resolve_ident(expr, scope) {
return value;
}
let segments: Vec<String> = path.path.segments.iter().map(|s| s.ident.to_string()).collect();
if segments.len() == 1 && segments[0] == "None" {
return DefaultValue::None;
}
if segments.len() >= 2
&& admits_enum_variant(field_ty)
&& let Some(name) = segments.last()
{
return DefaultValue::EnumVariant(name.clone());
}
unreadable(expr)
}
syn::Expr::Macro(mac) => {
let macro_name = mac
.mac
.path
.segments
.last()
.map(|s| s.ident.to_string())
.unwrap_or_default();
if !matches!(macro_name.as_str(), "vec" | "hashmap" | "hashset") {
return unreadable(expr);
}
if mac.mac.tokens.is_empty() {
return DefaultValue::Empty;
}
if macro_name != "vec" {
return unreadable(expr);
}
let Ok(elements) = mac
.mac
.parse_body_with(syn::punctuated::Punctuated::<syn::Expr, syn::Token![,]>::parse_terminated)
else {
return unreadable(expr);
};
if elements.is_empty() {
return DefaultValue::Empty;
}
let mut lowered = Vec::with_capacity(elements.len());
for element in &elements {
let value = expr_to_default_value(element, scope, field_ty);
if !carries_value(&value) {
return unreadable(expr);
}
lowered.push(value);
}
DefaultValue::ListLiteral(lowered)
}
_ => unreadable(expr),
}
}
fn resolve_ident(expr: &syn::Expr, scope: &EvalScope<'_>) -> Option<DefaultValue> {
let syn::Expr::Path(path) = expr else {
return None;
};
let segments: Vec<String> = path.path.segments.iter().map(|s| s.ident.to_string()).collect();
match segments.as_slice() {
[ident] => {
if let Some(value) = scope.params.get(ident) {
return Some(value.clone());
}
scope.literal_consts.get(ident).cloned()
}
[.., owner, name] => scope.associated_const(owner, name),
[] => None,
}
}
#[cfg(test)]
mod tests;