mod bridge_methods;
mod generator;
mod options_field;
mod registry;
mod visitor_bridge;
pub use crate::codegen::generators::trait_bridge::find_bridge_param;
pub use bridge_methods::gen_bridge_function;
pub use generator::Pyo3BridgeGenerator;
pub use options_field::gen_bridge_field_function;
pub use registry::{
collect_bridge_clear_fns, collect_bridge_register_fns, collect_bridge_unregister_fns, trait_bridge_imports,
};
use crate::codegen::generators::trait_bridge::{BridgeOutput, TraitBridgeSpec, gen_bridge_all};
use crate::core::config::TraitBridgeConfig;
use crate::core::ir::{ApiSurface, TypeDef};
use std::collections::{HashMap, HashSet};
use visitor_bridge::gen_visitor_bridge;
pub fn gen_trait_bridge(
trait_type: &TypeDef,
bridge_cfg: &TraitBridgeConfig,
core_import: &str,
error_type: &str,
error_constructor: &str,
api: &ApiSurface,
reexported_types: &[String],
) -> anyhow::Result<BridgeOutput> {
let type_paths: HashMap<String, String> = api
.types
.iter()
.map(|t| (t.name.clone(), t.rust_path.replace('-', "_")))
.chain(
api.enums
.iter()
.map(|e| (e.name.clone(), e.rust_path.replace('-', "_"))),
)
.chain(
api.excluded_type_paths
.iter()
.map(|(name, path)| (name.clone(), path.replace('-', "_"))),
)
.collect();
let is_visitor_bridge = bridge_cfg.type_alias.is_some()
&& bridge_cfg.register_fn.is_none()
&& bridge_cfg.super_trait.is_none()
&& trait_type.methods.iter().all(|m| m.has_default_impl);
if is_visitor_bridge {
let trait_path = trait_type.rust_path.replace('-', "_");
let struct_name = crate::codegen::generators::trait_bridge::bridge_wrapper_name("Py", bridge_cfg);
let code = gen_visitor_bridge(
trait_type,
bridge_cfg,
&struct_name,
&trait_path,
core_import,
&type_paths,
api,
)?;
Ok(BridgeOutput { imports: vec![], code })
} else {
let struct_param_types =
crate::codegen::generators::trait_bridge::native_marshalled_struct_params(trait_type, api);
let struct_return_types =
crate::codegen::generators::trait_bridge::native_marshalled_struct_returns(trait_type, api);
let forwardable_defaulted =
crate::codegen::generators::trait_bridge::forwardable_defaulted_method_names(trait_type, api);
let options_dataclass_types =
crate::backends::pyo3::gen_bindings::options_dataclass_type_names(api, reexported_types);
let unit_enum_return_types: HashSet<String> = api
.enums
.iter()
.filter(|e| e.variants.iter().all(|v| v.fields.is_empty()))
.map(|e| e.name.clone())
.collect();
let generator = Pyo3BridgeGenerator {
core_import: core_import.to_string(),
type_paths: type_paths.clone(),
error_type: error_type.to_string(),
struct_param_types,
struct_return_types,
forwardable_defaulted,
options_dataclass_types,
unit_enum_return_types,
};
let lifetime_type_names: HashSet<String> = api
.types
.iter()
.filter(|t| t.has_lifetime_params)
.map(|t| t.name.clone())
.collect();
let spec = TraitBridgeSpec {
trait_def: trait_type,
bridge_config: bridge_cfg,
core_import,
wrapper_prefix: "Py",
type_paths,
lifetime_type_names,
error_type: error_type.to_string(),
error_constructor: error_constructor.to_string(),
};
Ok(gen_bridge_all(&spec, &generator))
}
}
mod tests {
#[test]
fn trait_callback_runs_in_caller_contextvars_context() {
use crate::codegen::generators::trait_bridge::{TraitBridgeGenerator, TraitBridgeSpec};
use crate::core::config::TraitBridgeConfig;
use crate::core::ir::{MethodDef, ParamDef, ReceiverKind, TypeDef, TypeRef};
use std::collections::{HashMap, HashSet};
let trait_def = TypeDef {
name: "SampleService".to_owned(),
rust_path: "sample_core::SampleService".to_owned(),
is_trait: true,
is_opaque: true,
..TypeDef::default()
};
let bridge_cfg = TraitBridgeConfig {
trait_name: "SampleService".to_owned(),
register_fn: Some("register_sample".to_owned()),
registry_getter: Some("sample_core::registry::get".to_owned()),
..TraitBridgeConfig::default()
};
let spec = TraitBridgeSpec {
trait_def: &trait_def,
bridge_config: &bridge_cfg,
core_import: "sample_core",
wrapper_prefix: "Py",
type_paths: HashMap::new(),
lifetime_type_names: HashSet::new(),
error_type: "SampleError".to_owned(),
error_constructor: "SampleError::Message { message: {msg} }".to_owned(),
};
let generator = super::Pyo3BridgeGenerator {
core_import: "sample_core".to_owned(),
type_paths: HashMap::new(),
error_type: "SampleError".to_owned(),
struct_param_types: HashSet::new(),
struct_return_types: HashSet::new(),
forwardable_defaulted: HashSet::new(),
options_dataclass_types: HashSet::new(),
unit_enum_return_types: HashSet::new(),
};
let make_method = |is_async: bool| MethodDef {
name: "process".to_owned(),
params: vec![ParamDef {
name: "text".to_owned(),
ty: TypeRef::String,
..ParamDef::default()
}],
return_type: TypeRef::String,
is_async,
error_type: Some("SampleError".to_owned()),
receiver: Some(ReceiverKind::Ref),
..MethodDef::default()
};
let async_body = generator.gen_async_method_body(&make_method(true), &spec);
assert!(
async_body.contains("copy_context"),
"async bridge must capture the caller's contextvars context:\n{async_body}"
);
assert!(
async_body.contains("call_method1(\"run\""),
"async bridge must invoke the host method via ctx.run:\n{async_body}"
);
let sync_body = generator.gen_sync_method_body(&make_method(false), &spec);
assert!(
sync_body.contains("copy_context"),
"sync bridge must capture the caller's contextvars context:\n{sync_body}"
);
assert!(
sync_body.contains("call_method1(\"run\""),
"sync bridge must invoke the host method via ctx.run:\n{sync_body}"
);
}
#[test]
fn trait_callback_deserialize_error_names_return_type() {
use crate::codegen::generators::trait_bridge::{TraitBridgeGenerator, TraitBridgeSpec};
use crate::core::config::TraitBridgeConfig;
use crate::core::ir::{MethodDef, ReceiverKind, TypeDef, TypeRef};
use std::collections::{HashMap, HashSet};
let trait_def = TypeDef {
name: "SampleService".to_owned(),
rust_path: "sample_core::SampleService".to_owned(),
is_trait: true,
is_opaque: true,
..TypeDef::default()
};
let bridge_cfg = TraitBridgeConfig {
trait_name: "SampleService".to_owned(),
register_fn: Some("register_sample".to_owned()),
registry_getter: Some("sample_core::registry::get".to_owned()),
..TraitBridgeConfig::default()
};
let spec = TraitBridgeSpec {
trait_def: &trait_def,
bridge_config: &bridge_cfg,
core_import: "sample_core",
wrapper_prefix: "Py",
type_paths: HashMap::new(),
lifetime_type_names: HashSet::new(),
error_type: "SampleError".to_owned(),
error_constructor: "SampleError::Message { message: {msg} }".to_owned(),
};
let generator = super::Pyo3BridgeGenerator {
core_import: "sample_core".to_owned(),
type_paths: HashMap::new(),
error_type: "SampleError".to_owned(),
struct_param_types: HashSet::new(),
struct_return_types: HashSet::new(),
forwardable_defaulted: HashSet::new(),
options_dataclass_types: HashSet::new(),
unit_enum_return_types: HashSet::new(),
};
let make_method = |is_async: bool| MethodDef {
name: "build".to_owned(),
params: vec![],
return_type: TypeRef::Named("Doc".to_owned()),
is_async,
error_type: Some("SampleError".to_owned()),
receiver: Some(ReceiverKind::Ref),
..MethodDef::default()
};
for is_async in [true, false] {
let body = if is_async {
generator.gen_async_method_body(&make_method(true), &spec)
} else {
generator.gen_sync_method_body(&make_method(false), &spec)
};
assert!(
body.contains("expected return type `Doc`"),
"deserialize error must name the expected return type `Doc` (is_async={is_async}):\n{body}"
);
assert!(
body.contains("must be a mapping"),
"deserialize error must hint the value must be a mapping matching the type's fields (is_async={is_async}):\n{body}"
);
}
}
#[test]
fn trait_callback_native_struct_return_extracts_native_object_first() {
use crate::codegen::generators::trait_bridge::{TraitBridgeGenerator, TraitBridgeSpec};
use crate::core::config::TraitBridgeConfig;
use crate::core::ir::{MethodDef, ReceiverKind, TypeDef, TypeRef};
use std::collections::{HashMap, HashSet};
let trait_def = TypeDef {
name: "SampleService".to_owned(),
rust_path: "sample_core::SampleService".to_owned(),
is_trait: true,
is_opaque: true,
..TypeDef::default()
};
let bridge_cfg = TraitBridgeConfig {
trait_name: "SampleService".to_owned(),
register_fn: Some("register_sample".to_owned()),
registry_getter: Some("sample_core::registry::get".to_owned()),
..TraitBridgeConfig::default()
};
let spec = TraitBridgeSpec {
trait_def: &trait_def,
bridge_config: &bridge_cfg,
core_import: "sample_core",
wrapper_prefix: "Py",
type_paths: HashMap::new(),
lifetime_type_names: HashSet::new(),
error_type: "SampleError".to_owned(),
error_constructor: "SampleError::Message { message: {msg} }".to_owned(),
};
let generator = super::Pyo3BridgeGenerator {
core_import: "sample_core".to_owned(),
type_paths: HashMap::new(),
error_type: "SampleError".to_owned(),
struct_param_types: HashSet::new(),
struct_return_types: HashSet::from(["Doc".to_owned()]),
forwardable_defaulted: HashSet::new(),
options_dataclass_types: HashSet::new(),
unit_enum_return_types: HashSet::new(),
};
let make_method = |is_async: bool| MethodDef {
name: "build".to_owned(),
params: vec![],
return_type: TypeRef::Named("Doc".to_owned()),
is_async,
error_type: Some("SampleError".to_owned()),
receiver: Some(ReceiverKind::Ref),
..MethodDef::default()
};
for is_async in [true, false] {
let body = if is_async {
generator.gen_async_method_body(&make_method(true), &spec)
} else {
generator.gen_sync_method_body(&make_method(false), &spec)
};
assert!(
body.contains("extract::<Doc>()"),
"native return must try extracting the binding object `Doc` first (is_async={is_async}):\n{body}"
);
assert!(
body.contains("::from(native)"),
"native return must convert the extracted object via From<Binding> (is_async={is_async}):\n{body}"
);
assert!(
body.contains("serde_json::from_str"),
"the JSON/mapping fallback must remain (is_async={is_async}):\n{body}"
);
}
}
#[test]
fn trait_callback_owned_native_struct_param_is_marshalled() {
use crate::codegen::generators::trait_bridge::{TraitBridgeGenerator, TraitBridgeSpec};
use crate::core::config::TraitBridgeConfig;
use crate::core::ir::{MethodDef, ParamDef, ReceiverKind, TypeDef, TypeRef};
use std::collections::{HashMap, HashSet};
let trait_def = TypeDef {
name: "SampleExtractor".to_owned(),
rust_path: "sample_core::SampleExtractor".to_owned(),
is_trait: true,
is_opaque: true,
..TypeDef::default()
};
let bridge_cfg = TraitBridgeConfig {
trait_name: "SampleExtractor".to_owned(),
register_fn: Some("register_sample".to_owned()),
registry_getter: Some("sample_core::registry::get".to_owned()),
..TraitBridgeConfig::default()
};
let spec = TraitBridgeSpec {
trait_def: &trait_def,
bridge_config: &bridge_cfg,
core_import: "sample_core",
wrapper_prefix: "Py",
type_paths: HashMap::new(),
lifetime_type_names: HashSet::new(),
error_type: "SampleError".to_owned(),
error_constructor: "SampleError::Message { message: {msg} }".to_owned(),
};
let generator = super::Pyo3BridgeGenerator {
core_import: "sample_core".to_owned(),
type_paths: HashMap::new(),
error_type: "SampleError".to_owned(),
struct_param_types: HashSet::from(["Input".to_owned()]),
struct_return_types: HashSet::new(),
forwardable_defaulted: HashSet::new(),
options_dataclass_types: HashSet::new(),
unit_enum_return_types: HashSet::new(),
};
let make_method = |is_async: bool| MethodDef {
name: "handle".to_owned(),
params: vec![ParamDef {
name: "input".to_owned(),
ty: TypeRef::Named("Input".to_owned()),
is_ref: false,
..ParamDef::default()
}],
return_type: TypeRef::Unit,
is_async,
error_type: Some("SampleError".to_owned()),
receiver: Some(ReceiverKind::Ref),
..MethodDef::default()
};
for is_async in [true, false] {
let body = if is_async {
generator.gen_async_method_body(&make_method(true), &spec)
} else {
generator.gen_sync_method_body(&make_method(false), &spec)
};
assert!(
body.contains("Input::from("),
"owned native-struct param must be marshalled via From<core::T> (is_async={is_async}):\n{body}"
);
assert!(
!body.contains("(bound_method, input)") && !body.contains("(bound_method, input,"),
"owned native-struct param must not be handed to the host raw (is_async={is_async}):\n{body}"
);
}
}
#[test]
fn sync_unit_enum_return_accepts_bare_variant_name() {
use crate::codegen::generators::trait_bridge::{TraitBridgeGenerator, TraitBridgeSpec};
use crate::core::config::TraitBridgeConfig;
use crate::core::ir::{MethodDef, ReceiverKind, TypeDef, TypeRef};
use std::collections::{HashMap, HashSet};
let trait_def = TypeDef {
name: "SamplePostProcessor".to_owned(),
rust_path: "sample_core::SamplePostProcessor".to_owned(),
is_trait: true,
is_opaque: true,
..TypeDef::default()
};
let bridge_cfg = TraitBridgeConfig {
trait_name: "SamplePostProcessor".to_owned(),
register_fn: Some("register_sample".to_owned()),
registry_getter: Some("sample_core::registry::get".to_owned()),
..TraitBridgeConfig::default()
};
let spec = TraitBridgeSpec {
trait_def: &trait_def,
bridge_config: &bridge_cfg,
core_import: "sample_core",
wrapper_prefix: "Py",
type_paths: HashMap::new(),
lifetime_type_names: HashSet::new(),
error_type: "SampleError".to_owned(),
error_constructor: "SampleError::Message { message: {msg} }".to_owned(),
};
let generator = super::Pyo3BridgeGenerator {
core_import: "sample_core".to_owned(),
type_paths: HashMap::new(),
error_type: "SampleError".to_owned(),
struct_param_types: HashSet::new(),
struct_return_types: HashSet::new(),
forwardable_defaulted: HashSet::new(),
options_dataclass_types: HashSet::new(),
unit_enum_return_types: HashSet::from(["ProcessingStage".to_owned()]),
};
let method = MethodDef {
name: "processing_stage".to_owned(),
params: vec![],
return_type: TypeRef::Named("ProcessingStage".to_owned()),
is_async: false,
error_type: None,
receiver: Some(ReceiverKind::Ref),
..MethodDef::default()
};
let body = generator.gen_sync_method_body(&method, &spec);
assert!(
body.contains("serde_json::from_value(serde_json::Value::String"),
"unit-enum return must fall back to treating the string as a bare variant name:\n{body}"
);
assert!(
!body.contains("must be a mapping"),
"unit-enum deserialize error must not claim the value must be a mapping:\n{body}"
);
assert!(
body.contains("variant names"),
"unit-enum deserialize error should mention variant names:\n{body}"
);
}
#[test]
fn sync_struct_return_keeps_strict_mapping_error() {
use crate::codegen::generators::trait_bridge::{TraitBridgeGenerator, TraitBridgeSpec};
use crate::core::config::TraitBridgeConfig;
use crate::core::ir::{MethodDef, ReceiverKind, TypeDef, TypeRef};
use std::collections::{HashMap, HashSet};
let trait_def = TypeDef {
name: "SampleService".to_owned(),
rust_path: "sample_core::SampleService".to_owned(),
is_trait: true,
is_opaque: true,
..TypeDef::default()
};
let bridge_cfg = TraitBridgeConfig {
trait_name: "SampleService".to_owned(),
register_fn: Some("register_sample".to_owned()),
registry_getter: Some("sample_core::registry::get".to_owned()),
..TraitBridgeConfig::default()
};
let spec = TraitBridgeSpec {
trait_def: &trait_def,
bridge_config: &bridge_cfg,
core_import: "sample_core",
wrapper_prefix: "Py",
type_paths: HashMap::new(),
lifetime_type_names: HashSet::new(),
error_type: "SampleError".to_owned(),
error_constructor: "SampleError::Message { message: {msg} }".to_owned(),
};
let generator = super::Pyo3BridgeGenerator {
core_import: "sample_core".to_owned(),
type_paths: HashMap::new(),
error_type: "SampleError".to_owned(),
struct_param_types: HashSet::new(),
struct_return_types: HashSet::new(),
forwardable_defaulted: HashSet::new(),
options_dataclass_types: HashSet::new(),
unit_enum_return_types: HashSet::new(),
};
let method = MethodDef {
name: "build".to_owned(),
params: vec![],
return_type: TypeRef::Named("Doc".to_owned()),
is_async: false,
error_type: Some("SampleError".to_owned()),
receiver: Some(ReceiverKind::Ref),
..MethodDef::default()
};
let body = generator.gen_sync_method_body(&method, &spec);
assert!(
body.contains("must be a mapping"),
"struct deserialize error must keep the mapping-fields wording:\n{body}"
);
assert!(
!body.contains("serde_json::from_value(serde_json::Value::String"),
"struct return must not gain the unit-enum bare-string fallback:\n{body}"
);
}
#[test]
fn async_mut_param_writes_back_host_return_value() {
use crate::codegen::generators::trait_bridge::{TraitBridgeGenerator, TraitBridgeSpec};
use crate::core::config::TraitBridgeConfig;
use crate::core::ir::{MethodDef, ParamDef, ReceiverKind, TypeDef, TypeRef};
use std::collections::{HashMap, HashSet};
let trait_def = TypeDef {
name: "SamplePostProcessor".to_owned(),
rust_path: "sample_core::SamplePostProcessor".to_owned(),
is_trait: true,
is_opaque: true,
..TypeDef::default()
};
let bridge_cfg = TraitBridgeConfig {
trait_name: "SamplePostProcessor".to_owned(),
register_fn: Some("register_sample".to_owned()),
registry_getter: Some("sample_core::registry::get".to_owned()),
..TraitBridgeConfig::default()
};
let spec = TraitBridgeSpec {
trait_def: &trait_def,
bridge_config: &bridge_cfg,
core_import: "sample_core",
wrapper_prefix: "Py",
type_paths: HashMap::from([("Doc".to_owned(), "sample_core::Doc".to_owned())]),
lifetime_type_names: HashSet::new(),
error_type: "SampleError".to_owned(),
error_constructor: "SampleError::Message { message: {msg} }".to_owned(),
};
let generator = super::Pyo3BridgeGenerator {
core_import: "sample_core".to_owned(),
type_paths: HashMap::from([("Doc".to_owned(), "sample_core::Doc".to_owned())]),
error_type: "SampleError".to_owned(),
struct_param_types: HashSet::from(["Doc".to_owned()]),
struct_return_types: HashSet::new(),
forwardable_defaulted: HashSet::new(),
options_dataclass_types: HashSet::new(),
unit_enum_return_types: HashSet::new(),
};
let method = MethodDef {
name: "process".to_owned(),
params: vec![
ParamDef {
name: "result".to_owned(),
ty: TypeRef::Named("Doc".to_owned()),
is_ref: true,
is_mut: true,
..ParamDef::default()
},
ParamDef {
name: "config".to_owned(),
ty: TypeRef::String,
is_ref: true,
..ParamDef::default()
},
],
return_type: TypeRef::Unit,
is_async: true,
error_type: Some("SampleError".to_owned()),
receiver: Some(ReceiverKind::Ref),
..MethodDef::default()
};
let body = generator.gen_async_method_body(&method, &spec);
assert!(
body.contains("*result = "),
"the host's return value must be written back into the &mut param:\n{body}"
);
assert!(
!body.contains(".map(|_| ())"),
"the old discard-everything bridge shape must be gone for this method:\n{body}"
);
assert!(
body.contains("py_result.is_none()"),
"the host must be allowed to return None to mean \"unchanged\":\n{body}"
);
}
#[test]
fn visitor_bridge_uses_configured_context_and_result_metadata() {
let (api, trait_type, bridge) = crate::codegen::visitor_context::test_support::neutral_visitor_fixture();
let output = super::gen_trait_bridge(
&trait_type,
&bridge,
"sample_core",
"SampleError",
"SampleError::Message { message: {msg} }",
&api,
&[],
)
.expect("visitor bridge should generate");
crate::codegen::visitor_context::test_support::assert_neutral_visitor_output(&output.code);
assert!(output.code.contains("\"display_name\""));
}
}