use crate::core::ir::{ParamDef, TypeRef};
use ahash::AHashSet;
pub fn is_writeback_param(param: &ParamDef, opaque_types: &AHashSet<String>) -> bool {
if !param.is_ref || !param.is_mut || param.optional {
return false;
}
match ¶m.ty {
TypeRef::Named(name) => !opaque_types.contains(name.as_str()),
_ => false,
}
}
pub fn writeback_params<'a>(params: &'a [ParamDef], opaque_types: &AHashSet<String>) -> Vec<&'a ParamDef> {
params.iter().filter(|p| is_writeback_param(p, opaque_types)).collect()
}
pub fn writeback_param<'a>(
params: &'a [ParamDef],
return_type: &TypeRef,
opaque_types: &AHashSet<String>,
) -> Option<&'a ParamDef> {
if !matches!(return_type, TypeRef::Unit) {
return None;
}
let found = writeback_params(params, opaque_types);
match found.as_slice() {
[only] => Some(only),
_ => None,
}
}
pub fn writeback_type_name(param: &ParamDef) -> Option<&str> {
match ¶m.ty {
TypeRef::Named(name) => Some(name.as_str()),
_ => None,
}
}
pub fn effective_return_type(
params: &[ParamDef],
return_type: &TypeRef,
opaque_types: &AHashSet<String>,
) -> Option<TypeRef> {
writeback_param(params, return_type, opaque_types).map(|p| p.ty.clone())
}
pub fn reject_unsupported_writeback(
function_name: &str,
params: &[ParamDef],
return_type: &TypeRef,
opaque_types: &AHashSet<String>,
) -> anyhow::Result<()> {
let found = writeback_params(params, opaque_types);
if found.is_empty() {
return Ok(());
}
let names: Vec<&str> = found.iter().map(|p| p.name.as_str()).collect();
if found.len() > 1 {
anyhow::bail!(
"`{function_name}` takes {} `&mut` parameters ({}); a binding has one return slot, so at \
most one is supported. Change the core signature to take one `&mut` parameter, or to take \
the values by move and return them.",
found.len(),
names.join(", "),
);
}
if !matches!(return_type, TypeRef::Unit) {
anyhow::bail!(
"`{function_name}` takes a `&mut` parameter (`{}`) and also returns a value; the single \
return slot already carries the updated `&mut` value. Change the core signature to return \
the updated value itself, or to fold both results into one returned type.",
names[0],
);
}
Ok(())
}
#[cfg(test)]
mod tests;