use crate::core::config::TraitBridgeConfig;
use crate::core::ir::ApiSurface;
pub(super) const BRIDGE_CTOR_PARAMS: &str = "rb_obj: magnus::Value, name: String";
pub(super) const BRIDGE_CTOR_RETURN_TYPE: &str = "Result<Self, magnus::Error>";
pub(super) const BRIDGE_CTOR_IS_FALLIBLE: bool = true;
fn named_type_name(ty: &crate::core::ir::TypeRef) -> Option<&str> {
match ty {
crate::core::ir::TypeRef::Named(n) => Some(n.as_str()),
crate::core::ir::TypeRef::Optional(inner) => match inner.as_ref() {
crate::core::ir::TypeRef::Named(n) => Some(n.as_str()),
_ => None,
},
_ => None,
}
}
fn bridge_name_expr(func: &crate::core::ir::FunctionDef, bridge_param_idx: usize) -> &'static str {
let has_name = func.params.iter().enumerate().any(|(idx, param)| {
idx != bridge_param_idx
&& param.name == "name"
&& !param.optional
&& matches!(param.ty, crate::core::ir::TypeRef::String)
});
if has_name { "name.clone()" } else { "String::new()" }
}
#[allow(clippy::too_many_arguments)]
pub fn gen_bridge_function(
api: &ApiSurface,
func: &crate::core::ir::FunctionDef,
bridge_param_idx: usize,
bridge_cfg: &TraitBridgeConfig,
mapper: &dyn crate::codegen::type_mapper::TypeMapper,
opaque_types: &ahash::AHashSet<String>,
default_types: &std::collections::HashSet<&str>,
core_import: &str,
) -> String {
use crate::core::ir::TypeRef;
let struct_name = crate::codegen::generators::trait_bridge::bridge_wrapper_name("Rb", bridge_cfg);
let handle_path = crate::codegen::generators::trait_bridge::bridge_handle_path(api, bridge_cfg, core_import);
let param_name = &func.params[bridge_param_idx].name;
let bridge_param = &func.params[bridge_param_idx];
let is_optional = bridge_param.optional || matches!(&bridge_param.ty, TypeRef::Optional(_));
let bridge_name = bridge_name_expr(func, bridge_param_idx);
let (bridge_handle_type, bridge_value) = if bridge_param.core_wrapper == crate::core::ir::CoreWrapper::Arc {
(
format!("std::sync::Arc<dyn {handle_path}>"),
"std::sync::Arc::new(bridge)".to_string(),
)
} else {
(
handle_path.clone(),
"std::sync::Arc::new(std::sync::Mutex::new(bridge))".to_string(),
)
};
let mut sig_parts = Vec::new();
for (idx, p) in func.params.iter().enumerate() {
if idx == bridge_param_idx {
if is_optional {
sig_parts.push(format!("{}: Option<magnus::Value>", p.name));
} else {
sig_parts.push(format!("{}: magnus::Value", p.name));
}
} else {
let promoted = (is_optional && idx > bridge_param_idx) || func.params[..idx].iter().any(|pp| pp.optional);
let ty = if p.optional || promoted {
format!("Option<{}>", mapper.map_type(&p.ty))
} else {
mapper.map_type(&p.ty)
};
sig_parts.push(format!("{}: {}", p.name, ty));
}
}
let params_str = sig_parts.join(", ");
let return_type = mapper.map_type(&func.return_type);
let deser_params: Vec<_> = func
.params
.iter()
.enumerate()
.filter(|(idx, p)| {
*idx != bridge_param_idx && named_type_name(&p.ty).is_some_and(|n| !opaque_types.contains(n))
})
.collect();
let params_need_fallible_deser = deser_params
.iter()
.copied()
.any(|(_, p)| !default_types.contains(named_type_name(&p.ty).unwrap_or_default()));
let ctor_try = if BRIDGE_CTOR_IS_FALLIBLE { "?" } else { "" };
let bridge_wrap = if is_optional {
format!(
"let {param_name}: Option<{bridge_handle_type}> = match {param_name} {{\n \
Some(v) if !v.is_nil() => {{\n \
let bridge = {struct_name}::new(v, {bridge_name}){ctor_try};\n \
Some({bridge_value} as {bridge_handle_type})\n \
}},\n \
_ => None,\n \
}};"
)
} else {
format!(
"let {param_name} = {{\n \
let bridge = {struct_name}::new({param_name}, {bridge_name}){ctor_try};\n \
{bridge_value} as {bridge_handle_type}\n \
}};"
)
};
let bridge_construction_is_fallible = bridge_wrap.contains('?');
let has_error = func.error_type.is_some() || params_need_fallible_deser || bridge_construction_is_fallible;
let ret = mapper.wrap_return(&return_type, has_error);
let err_conv = ".map_err(|e| magnus::Error::new(unsafe { magnus::Ruby::get_unchecked() }.exception_runtime_error(), e.to_string()))";
let serde_bindings: String = deser_params
.into_iter()
.map(|(_, p)| {
let name = &p.name;
let named_type = named_type_name(&p.ty).unwrap_or_default().to_string();
let core_path = format!("{core_import}::{named_type}");
let is_dt = default_types.contains(named_type.as_str());
if is_dt {
if p.optional || matches!(&p.ty, TypeRef::Optional(_)) {
format!("let {name}_core: Option<{core_path}> = {name}.map(Into::into);\n ")
} else {
format!("let {name}_core: {core_path} = {name}.into();\n ")
}
} else if p.optional || matches!(&p.ty, TypeRef::Optional(_)) {
format!(
"let {name}_core: Option<{core_path}> = {name}.as_deref().filter(|s| *s != \"nil\").map(|s| serde_json::from_str(s){err_conv}).transpose()?;\n "
)
} else {
format!(
"let {name}_core: {core_path} = serde_json::from_str(&{name}){err_conv}?;\n "
)
}
})
.collect();
let call_args: Vec<String> = func
.params
.iter()
.enumerate()
.map(|(idx, p)| {
if idx == bridge_param_idx {
return p.name.clone();
}
match &p.ty {
TypeRef::Named(n) if opaque_types.contains(n.as_str()) => {
if p.optional {
format!("{}.as_ref().map(|v| &v.inner)", p.name)
} else {
format!("&{}.inner", p.name)
}
}
TypeRef::Named(_) => format!("{}_core", p.name),
TypeRef::Optional(inner) => {
if let TypeRef::Named(n) = inner.as_ref() {
if opaque_types.contains(n.as_str()) {
format!("{}.as_ref().map(|v| &v.inner)", p.name)
} else {
format!("{}_core", p.name)
}
} else {
p.name.clone()
}
}
TypeRef::String | TypeRef::Char => {
if p.is_ref {
format!("&{}", p.name)
} else {
p.name.clone()
}
}
_ => p.name.clone(),
}
})
.collect();
let call_args_str = call_args.join(", ");
let core_fn_path = {
let path = func.rust_path.replace('-', "_");
if path.starts_with(core_import) {
path
} else {
format!("{core_import}::{}", func.name)
}
};
let core_call = format!("{core_fn_path}({call_args_str})");
let return_wrap = match &func.return_type {
TypeRef::Named(name) if opaque_types.contains(name.as_str()) => {
format!("{name} {{ inner: std::sync::Arc::new(val) }}")
}
TypeRef::Named(_) => "val.into()".to_string(),
TypeRef::String | TypeRef::Bytes => "val.into()".to_string(),
_ => "val".to_string(),
};
let body = if func.error_type.is_some() {
if return_wrap == "val" {
format!("{bridge_wrap}\n {serde_bindings}{core_call}{err_conv}")
} else {
format!("{bridge_wrap}\n {serde_bindings}{core_call}.map(|val| {return_wrap}){err_conv}")
}
} else {
if return_wrap == "val" {
format!("{bridge_wrap}\n {serde_bindings}Ok({core_call})")
} else {
format!("{bridge_wrap}\n {serde_bindings}let val = {core_call};\n Ok({return_wrap})")
}
};
let func_name = &func.name;
let mut out = String::with_capacity(1024);
if func.error_type.is_some() {
out.push_str("#[allow(clippy::missing_errors_doc)]\n");
}
out.push_str("#[allow(unused_variables)]\n");
let sig = format!("pub fn {func_name}({params_str}) -> {ret} {{\n {body}\n}}\n");
out.push_str(&sig);
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backends::magnus::type_map::MagnusMapper;
use crate::core::ir::{FunctionDef, ParamDef, TypeRef};
fn arity_after(text: &str, needle: &str) -> usize {
let start = text
.find(needle)
.unwrap_or_else(|| panic!("{needle:?} not found in:\n{text}"))
+ needle.len();
let mut depth = 1usize;
let mut top_level_commas = 0usize;
for ch in text[start..].chars() {
match ch {
'(' => depth += 1,
')' => {
depth -= 1;
if depth == 0 {
break;
}
}
',' if depth == 1 => top_level_commas += 1,
_ => {}
}
}
top_level_commas + 1
}
fn close_paren_after(text: &str, needle: &str) -> usize {
let start = text
.find(needle)
.unwrap_or_else(|| panic!("{needle:?} not found in:\n{text}"))
+ needle.len();
let mut depth = 1usize;
for (i, ch) in text[start..].char_indices() {
match ch {
'(' => depth += 1,
')' => {
depth -= 1;
if depth == 0 {
return start + i + 1;
}
}
_ => {}
}
}
panic!("unbalanced parens after {needle:?} in:\n{text}");
}
fn call_site_is_fallible(text: &str, needle: &str) -> bool {
text[close_paren_after(text, needle)..].starts_with('?')
}
fn definition_is_fallible(text: &str, needle: &str) -> bool {
let after_close = close_paren_after(text, needle);
let brace = text[after_close..]
.find('{')
.unwrap_or_else(|| panic!("no constructor body opening brace after {needle:?} in:\n{text}"));
text[after_close..after_close + brace].contains("Result")
}
#[test]
fn visitor_bridge_constructor_matches_call_site_arity_and_fallibility() {
let (api, trait_type, bridge_cfg) = crate::codegen::visitor_context::test_support::neutral_visitor_fixture();
let definition = super::super::gen_trait_bridge(
&trait_type,
&bridge_cfg,
"sample_core",
"SampleError",
"SampleError::Message { message: {msg} }",
&api,
)
.expect("visitor bridge definition should generate");
assert!(
definition.contains("pub fn new("),
"visitor bridge must declare a constructor:\n{definition}"
);
let def_arity = arity_after(&definition, "pub fn new(");
let def_is_fallible = definition_is_fallible(&definition, "pub fn new(");
let func = FunctionDef {
name: "inspect".to_string(),
rust_path: "sample_core::inspect".to_string(),
params: vec![ParamDef {
name: "walker".to_string(),
ty: TypeRef::Named("DocumentWalkerHandle".to_string()),
optional: true,
..ParamDef::default()
}],
return_type: TypeRef::Unit,
..FunctionDef::default()
};
let call_site = gen_bridge_function(
&api,
&func,
0,
&bridge_cfg,
&MagnusMapper,
&ahash::AHashSet::default(),
&std::collections::HashSet::new(),
"sample_core",
);
let struct_name = crate::codegen::generators::trait_bridge::bridge_wrapper_name("Rb", &bridge_cfg);
let ctor_call_needle = format!("{struct_name}::new(");
assert!(
call_site.contains(&ctor_call_needle),
"call site must construct the bridge via {ctor_call_needle:?}:\n{call_site}"
);
let call_arity = arity_after(&call_site, &ctor_call_needle);
let call_fallible = call_site_is_fallible(&call_site, &ctor_call_needle);
assert_eq!(
def_arity, call_arity,
"constructor definition arity ({def_arity}) must match call site arity ({call_arity}); \
definition:\n{definition}\ncall site:\n{call_site}"
);
assert_eq!(
def_is_fallible, call_fallible,
"constructor definition fallibility ({def_is_fallible}) must match call site \
fallibility ({call_fallible}); definition:\n{definition}\ncall site:\n{call_site}"
);
}
}