use super::{DefaultValue, EvalScope, TypeRef, carries_value, expr_to_default_value, struct_expr_defaults};
use ahash::AHashMap;
use quote::ToTokens;
pub(super) struct StructBody<'a> {
struct_expr: &'a syn::ExprStruct,
mutations: Vec<FieldMutation<'a>>,
}
struct FieldMutation<'a> {
field: String,
kind: MutationKind<'a>,
source: String,
}
enum MutationKind<'a> {
Assign(&'a syn::Expr),
Push(&'a syn::Expr),
Extend(&'a syn::Expr),
Opaque,
}
pub(super) fn read_struct_body(block: &syn::Block) -> Option<StructBody<'_>> {
if let Some(struct_expr) = tail_struct_expr(block) {
return Some(StructBody {
struct_expr,
mutations: Vec::new(),
});
}
read_mutated_body(block)
}
pub(super) fn struct_body_defaults(body: &StructBody<'_>, scope: &EvalScope<'_>) -> AHashMap<String, DefaultValue> {
let mut defaults = struct_expr_defaults(body.struct_expr, scope);
for mutation in &body.mutations {
let field_ty = scope.field_types.get(&mutation.field);
match &mutation.kind {
MutationKind::Assign(value) => {
defaults.insert(mutation.field.clone(), expr_to_default_value(value, scope, field_ty));
}
MutationKind::Push(value) => {
let current = defaults
.entry(mutation.field.clone())
.or_insert_with(|| DefaultValue::Unresolved(mutation.source.clone()));
push(current, value, scope, field_ty, &mutation.source);
}
MutationKind::Extend(value) => {
let current = defaults
.entry(mutation.field.clone())
.or_insert_with(|| DefaultValue::Unresolved(mutation.source.clone()));
extend(current, value, scope, field_ty, &mutation.source);
}
MutationKind::Opaque => {
defaults.insert(
mutation.field.clone(),
DefaultValue::Unresolved(mutation.source.clone()),
);
}
}
}
defaults
}
fn push(
current: &mut DefaultValue,
value: &syn::Expr,
scope: &EvalScope<'_>,
field_ty: Option<&TypeRef>,
source: &str,
) {
let Some(TypeRef::Vec(element_ty)) = field_ty else {
*current = DefaultValue::Unresolved(source.to_string());
return;
};
let element = expr_to_default_value(value, scope, Some(element_ty));
if !carries_value(&element) {
*current = DefaultValue::Unresolved(source.to_string());
return;
}
match current {
DefaultValue::Empty => *current = DefaultValue::ListLiteral(vec![element]),
DefaultValue::ListLiteral(elements) => elements.push(element),
_ => *current = DefaultValue::Unresolved(source.to_string()),
}
}
fn extend(
current: &mut DefaultValue,
value: &syn::Expr,
scope: &EvalScope<'_>,
field_ty: Option<&TypeRef>,
source: &str,
) {
let Some(TypeRef::Vec(element_ty)) = field_ty else {
*current = DefaultValue::Unresolved(source.to_string());
return;
};
let addition = expr_to_default_value(value, scope, Some(element_ty));
let additions = match addition {
DefaultValue::Empty => Vec::new(),
DefaultValue::ListLiteral(elements) => elements,
_ => {
*current = DefaultValue::Unresolved(source.to_string());
return;
}
};
match current {
DefaultValue::Empty if additions.is_empty() => {}
DefaultValue::Empty => *current = DefaultValue::ListLiteral(additions),
DefaultValue::ListLiteral(elements) => elements.extend(additions),
_ => *current = DefaultValue::Unresolved(source.to_string()),
}
}
fn tail_struct_expr(block: &syn::Block) -> Option<&syn::ExprStruct> {
let (tail, leading) = block.stmts.split_last()?;
if is_attributed(tail) {
return None;
}
if leading
.iter()
.any(|stmt| is_attributed(stmt) || contains_early_return(stmt) || contains_macro(stmt))
{
return None;
}
let syn::Stmt::Expr(expr, _) = tail else {
return None;
};
unwrap_to_struct_expr(expr)
}
fn unwrap_to_struct_expr(expr: &syn::Expr) -> Option<&syn::ExprStruct> {
match expr {
syn::Expr::Struct(s) if s.attrs.is_empty() => Some(s),
syn::Expr::Block(b) if b.attrs.is_empty() => tail_struct_expr(&b.block),
_ => None,
}
}
fn read_mutated_body(block: &syn::Block) -> Option<StructBody<'_>> {
let (tail, leading) = block.stmts.split_last()?;
if is_attributed(tail) {
return None;
}
if leading
.iter()
.any(|stmt| is_attributed(stmt) || contains_early_return(stmt))
{
return None;
}
let [first, mutation_stmts @ ..] = leading else {
return None;
};
let (binding, struct_expr) = local_struct_binding(first)?;
if !tail_returns_binding(tail, &binding) {
return None;
}
let mut mutations = Vec::with_capacity(mutation_stmts.len());
for stmt in mutation_stmts {
mutations.push(classify_mutation(stmt, &binding)?);
}
Some(StructBody { struct_expr, mutations })
}
fn local_struct_binding(stmt: &syn::Stmt) -> Option<(String, &syn::ExprStruct)> {
let syn::Stmt::Local(local) = stmt else {
return None;
};
let init = local.init.as_ref()?;
if init.diverge.is_some() {
return None;
}
let syn::Expr::Struct(struct_expr) = init.expr.as_ref() else {
return None;
};
if !struct_expr.attrs.is_empty() {
return None;
}
if struct_expr.rest.is_some() {
return None;
}
Some((binding_ident(&local.pat)?, struct_expr))
}
fn binding_ident(pat: &syn::Pat) -> Option<String> {
match pat {
syn::Pat::Ident(pat_ident) if pat_ident.by_ref.is_none() && pat_ident.subpat.is_none() => {
Some(pat_ident.ident.to_string())
}
syn::Pat::Type(pat_type) => binding_ident(&pat_type.pat),
_ => None,
}
}
fn tail_returns_binding(stmt: &syn::Stmt, binding: &str) -> bool {
let syn::Stmt::Expr(expr, _) = stmt else {
return false;
};
let returned = match expr {
syn::Expr::Return(ret) => match ret.expr.as_deref() {
Some(inner) => inner,
None => return false,
},
other => other,
};
matches!(returned, syn::Expr::Path(path) if path.qself.is_none() && path.path.is_ident(binding))
}
fn classify_mutation<'a>(stmt: &'a syn::Stmt, binding: &str) -> Option<FieldMutation<'a>> {
let syn::Stmt::Expr(expr, Some(_)) = stmt else {
return None;
};
let source = expr.to_token_stream().to_string();
match expr {
syn::Expr::Assign(assign) => {
let field = binding_field(&assign.left, binding)?;
reject_escape(&assign.right, binding)?;
Some(FieldMutation {
field,
kind: MutationKind::Assign(&assign.right),
source,
})
}
syn::Expr::MethodCall(call) => {
let field = binding_field(&call.receiver, binding)?;
for argument in &call.args {
reject_escape(argument, binding)?;
}
let arguments: Vec<&syn::Expr> = call.args.iter().collect();
let kind = match (call.method.to_string().as_str(), arguments.as_slice()) {
("push", [value]) => MutationKind::Push(value),
("extend", [value]) => MutationKind::Extend(value),
("insert", _) => MutationKind::Opaque,
_ => return None,
};
Some(FieldMutation { field, kind, source })
}
_ => None,
}
}
fn is_attributed(stmt: &syn::Stmt) -> bool {
stmt.to_token_stream().to_string().trim_start().starts_with('#')
}
fn contains_early_return(stmt: &syn::Stmt) -> bool {
struct Scan {
found: bool,
}
impl<'ast> syn::visit::Visit<'ast> for Scan {
fn visit_expr_return(&mut self, _node: &'ast syn::ExprReturn) {
self.found = true;
}
fn visit_expr_closure(&mut self, _node: &'ast syn::ExprClosure) {}
}
let mut scan = Scan { found: false };
syn::visit::Visit::visit_stmt(&mut scan, stmt);
scan.found
}
fn contains_macro(stmt: &syn::Stmt) -> bool {
struct Scan {
found: bool,
}
impl<'ast> syn::visit::Visit<'ast> for Scan {
fn visit_macro(&mut self, _node: &'ast syn::Macro) {
self.found = true;
}
}
let mut scan = Scan { found: false };
syn::visit::Visit::visit_stmt(&mut scan, stmt);
scan.found
}
fn binding_field(expr: &syn::Expr, binding: &str) -> Option<String> {
let syn::Expr::Field(field) = expr else {
return None;
};
let syn::Expr::Path(path) = field.base.as_ref() else {
return None;
};
if path.qself.is_some() || !path.path.is_ident(binding) {
return None;
}
match &field.member {
syn::Member::Named(ident) => Some(ident.to_string()),
syn::Member::Unnamed(_) => None,
}
}
fn reject_escape(expr: &syn::Expr, binding: &str) -> Option<()> {
(!mentions_binding(expr, binding)).then_some(())
}
fn mentions_binding(expr: &syn::Expr, binding: &str) -> bool {
expr.to_token_stream()
.to_string()
.split(|c: char| !c.is_alphanumeric() && c != '_')
.any(|token| token == binding)
}