This is an automated email from the ASF dual-hosted git repository. tlopex pushed a commit to branch rust-mutate-context-dispatch in repository https://gitbox.apache.org/repos/asf/tvm-ffi.git
commit ba9883ec06c24a95d256ef795efcef2412e0d9b7 Author: tlopex <[email protected]> AuthorDate: Mon Aug 31 20:52:18 2026 -0400 [FEAT][RUST] Pass mutation context to typed dispatch callbacks --- docs/guides/rust_lang_guide.md | 32 +++-- rust/tvm-ffi-macros/src/dispatch.rs | 182 +++++++++++++++++++------- rust/tvm-ffi/src/extra/structural_mutate.rs | 56 +++++++- rust/tvm-ffi/src/lib.rs | 2 +- rust/tvm-ffi/tests/test_structural_mutate.rs | 188 ++++++++++++++++++++------- 5 files changed, 351 insertions(+), 109 deletions(-) diff --git a/docs/guides/rust_lang_guide.md b/docs/guides/rust_lang_guide.md index 92f2d059..67ebcbbf 100644 --- a/docs/guides/rust_lang_guide.md +++ b/docs/guides/rust_lang_guide.md @@ -529,37 +529,49 @@ assert_eq!(mutator.state().integers, 2); `maybe_inplace_mutate` preserves the reuse opportunity of an owned value. Callbacks are `Fn`; mutable data belongs in the mutator state. -For ordinary mutable state, `#[dispatch(mutate)]` generates a -`StructuralMutator` from `mutate_*` methods. A matching handler returns the -current value's final result and may recursively call `self.mutate()`; -an unmatched value follows default mutation with its current in-place permit: +`#[dispatch(mutate)]` groups typed `mutate_*` callbacks. Dispatch only selects +the first matching callback; `MutateContext` supplies recursion, the current +definition region, and mutable state. `mutator.mutate(child)` inherits the +current region, while `mutate_with` is available for an explicit override. An +unmatched value follows default mutation with its current in-place permit: ```rust -use tvm_ffi::{dispatch, structural_mutate, Any, Array, DefRegionKind}; +use tvm_ffi::{ + dispatch, structural_mutate, Any, Array, MutateCallbacks, MutateContext, +}; #[derive(Default)] -struct Increment { +struct IncrementState { integers: usize, } +struct Increment; + #[dispatch(mutate)] impl Increment { - fn mutate_integer(&mut self, value: i64, _kind: DefRegionKind) -> Any { - self.integers += 1; + fn mutate_integer( + &self, + value: i64, + mutator: &mut MutateContext<'_, IncrementState>, + ) -> Any { + mutator.state_mut().integers += 1; Any::from(value + 1) } } -let mut increment = Increment::default(); +let mut increment = MutateCallbacks::new(IncrementState::default(), Increment); let mutated = structural_mutate( Array::new(vec![1_i64, 2]), &mut increment, )?; let mutated = Array::<i64>::try_from(mutated)?; assert_eq!(mutated.iter().collect::<Vec<_>>(), vec![2, 3]); -assert_eq!(increment.integers, 2); +assert_eq!(increment.state().integers, 2); ``` +A generated dispatch whose context state is `()` can be passed directly to +`structural_mutate`, without `MutateCallbacks`. + For a named custom recursion policy, implement `StructuralMutator` and pass `&mut` it to `structural_mutate`. `InplaceValue` is an engine-issued capability: callers cannot construct it from a read-only `MapValue`. Override diff --git a/rust/tvm-ffi-macros/src/dispatch.rs b/rust/tvm-ffi-macros/src/dispatch.rs index 21c6ccb4..5ba37e39 100644 --- a/rust/tvm-ffi-macros/src/dispatch.rs +++ b/rust/tvm-ffi-macros/src/dispatch.rs @@ -20,7 +20,10 @@ use proc_macro::TokenStream; use proc_macro2::TokenStream as TokenStream2; use quote::{quote, quote_spanned}; -use syn::{parse_macro_input, FnArg, ImplItem, ImplItemMethod, ItemImpl, Meta, NestedMeta, Type}; +use syn::{ + parse_macro_input, FnArg, GenericArgument, ImplItem, ImplItemMethod, ItemImpl, Meta, + NestedMeta, PathArguments, Type, +}; use crate::utils::get_tvm_ffi_crate; @@ -77,8 +80,8 @@ impl DispatchMode { fn result_is_optional(self) -> bool { match self { - Self::Walk | Self::Map => true, - Self::Visit | Self::Mutate => false, + Self::Walk | Self::Map | Self::Mutate => true, + Self::Visit => false, } } } @@ -102,11 +105,18 @@ impl syn::parse::Parse for DispatchArgs { )); }; if !input.is_empty() { - return Err(input.error(format!( - "`dispatch({})` takes no further arguments; a handler that needs the \ - definition-region state declares a trailing `DefRegionKind` argument", - mode.name() - ))); + let message = if matches!(mode, DispatchMode::Mutate) { + "`dispatch(mutate)` takes no further arguments; the definition region is \ + available through `MutateContext::region()`" + .to_owned() + } else { + format!( + "`dispatch({})` takes no further arguments; a handler that needs the \ + definition-region state declares a trailing `DefRegionKind` argument", + mode.name() + ) + }; + return Err(input.error(message)); } Ok(DispatchArgs { mode }) } @@ -159,6 +169,16 @@ fn expand(item_impl: &ItemImpl, mode: DispatchMode) -> syn::Result<TokenStream2> #tvm_ffi::extra::structural_mutate::IntoMutateResult::into_mutate_result }, }; + let mutate_state = if matches!(mode, DispatchMode::Mutate) { + Some( + handlers[0] + .mutate_state + .as_ref() + .expect("mutate handlers always record their context state"), + ) + } else { + None + }; let links = expand_links(&handlers, mode, &into_result, quote!(value)); let self_type = &item_impl.self_ty; let (impl_generics, _, where_clause) = item_impl.generics.split_for_impl(); @@ -236,34 +256,25 @@ fn expand(item_impl: &ItemImpl, mode: DispatchMode) -> syn::Result<TokenStream2> } }, DispatchMode::Mutate => { - let inplace_links = - expand_links(&handlers, mode, &into_result, quote!(value.as_value())); + let state = mutate_state.expect("mutate dispatch has a context state"); quote! { - impl #impl_generics #tvm_ffi::extra::structural_mutate::StructuralMutator + impl #impl_generics #tvm_ffi::extra::structural_mutate::MutateDispatch for #self_type #where_clause { + type State = #state; + #[inline] #[allow(unreachable_code, unused_variables)] fn dispatch_mutate( - &mut self, + &self, value: &#tvm_ffi::extra::structural_mutate::MapValue, - def_region_kind: #tvm_ffi::extra::structural_visit::DefRegionKind, - ) -> #tvm_ffi::Result<#tvm_ffi::Any> { + mutator: &mut #tvm_ffi::extra::structural_mutate::MutateContext< + '_, + Self::State, + >, + ) -> Option<#tvm_ffi::extra::structural_mutate::MutateResult> { #(#links)* - <Self as #tvm_ffi::extra::structural_mutate::StructuralMutator>:: - default_mutate(self, value, def_region_kind) - } - - #[inline] - #[allow(unreachable_code, unused_variables)] - fn dispatch_maybe_inplace_mutate( - &mut self, - value: #tvm_ffi::extra::structural_mutate::InplaceValue<'_>, - def_region_kind: #tvm_ffi::extra::structural_visit::DefRegionKind, - ) -> #tvm_ffi::Result<#tvm_ffi::Any> { - #(#inplace_links)* - <Self as #tvm_ffi::extra::structural_mutate::StructuralMutator>:: - default_maybe_inplace_mutate(self, value, def_region_kind) + None } } } @@ -289,10 +300,10 @@ fn expand_links( .map(|handler| { let method = &handler.method; let attrs = &handler.cfg_attrs; - let kind_arg = if handler.wants_def_region { - quote!(, def_region_kind) - } else { - quote!() + let trailing_arg = match mode { + DispatchMode::Mutate => quote!(, mutator), + _ if handler.wants_def_region => quote!(, def_region_kind), + _ => quote!(), }; let wrap_result = |result: TokenStream2| { if mode.result_is_optional() { @@ -304,7 +315,7 @@ fn expand_links( let invoke = match &handler.argument { HandlerArgument::Value => { let result = wrap_result(quote! { - #into_result(self.#method(#value #kind_arg)) + #into_result(self.#method(#value #trailing_arg)) }); quote! { return #result; @@ -312,7 +323,7 @@ fn expand_links( } HandlerArgument::BorrowedNode(node_type) => { let result = wrap_result(quote! { - #into_result(self.#method(node #kind_arg)) + #into_result(self.#method(node #trailing_arg)) }); quote! { if let Some(node) = #value.as_node::<#node_type>() { @@ -322,7 +333,7 @@ fn expand_links( } HandlerArgument::Owned(value_type) => { let result = wrap_result(quote! { - #into_result(self.#method(typed #kind_arg)) + #into_result(self.#method(typed #trailing_arg)) }); quote! { if let Some(typed) = #value.cast::<#value_type>() { @@ -345,6 +356,7 @@ struct Handler { method: syn::Ident, argument: HandlerArgument, wants_def_region: bool, + mutate_state: Option<Type>, cfg_attrs: Vec<Meta>, } @@ -356,22 +368,43 @@ enum HandlerArgument { fn parse_handler(method: &ImplItemMethod, mode: DispatchMode) -> syn::Result<Handler> { let inputs = &method.sig.inputs; - let receiver_is_mut = matches!( - inputs.first(), - Some(FnArg::Receiver(receiver)) - if receiver.reference.is_some() && receiver.mutability.is_some() - ); - if !receiver_is_mut || !(inputs.len() == 2 || inputs.len() == 3) { - return Err(syn::Error::new_spanned( - &method.sig, + let receiver_is_expected = match (mode, inputs.first()) { + (DispatchMode::Mutate, Some(FnArg::Receiver(receiver))) => { + receiver.reference.is_some() && receiver.mutability.is_none() + } + (_, Some(FnArg::Receiver(receiver))) => { + receiver.reference.is_some() && receiver.mutability.is_some() + } + _ => false, + }; + let arity_is_expected = if matches!(mode, DispatchMode::Mutate) { + inputs.len() == 3 + } else { + inputs.len() == 2 || inputs.len() == 3 + }; + if !receiver_is_expected || !arity_is_expected { + let message = if matches!(mode, DispatchMode::Mutate) { + "mutate handlers must take `&self`, a node, and `&mut MutateContext<'_, State>`" + .to_owned() + } else { format!( "{} handlers must take `&mut self`, a node, and optionally a trailing \ `DefRegionKind` argument", mode.name() - ), - )); + ) + }; + return Err(syn::Error::new_spanned(&method.sig, message)); } - let wants_def_region = inputs.len() == 3; + let wants_def_region = !matches!(mode, DispatchMode::Mutate) && inputs.len() == 3; + let mutate_state = if matches!(mode, DispatchMode::Mutate) { + let context_type = match inputs.iter().nth(2) { + Some(FnArg::Typed(context)) => context.ty.as_ref(), + _ => unreachable!("the third argument cannot be a receiver"), + }; + Some(parse_mutate_context_state(context_type)?) + } else { + None + }; let value_type = match inputs.iter().nth(1) { Some(FnArg::Typed(value)) => (*value.ty).clone(), @@ -401,10 +434,67 @@ fn parse_handler(method: &ImplItemMethod, mode: DispatchMode) -> syn::Result<Han method: method.sig.ident.clone(), argument, wants_def_region, + mutate_state, cfg_attrs, }) } +fn parse_mutate_context_state(context_type: &Type) -> syn::Result<Type> { + let Type::Reference(reference) = context_type else { + return Err(syn::Error::new_spanned( + context_type, + "the mutate context must be `&mut MutateContext<'_, State>`", + )); + }; + if reference.mutability.is_none() { + return Err(syn::Error::new_spanned( + context_type, + "the mutate context must be a mutable reference", + )); + } + let Type::Path(path) = reference.elem.as_ref() else { + return Err(syn::Error::new_spanned( + context_type, + "expected `&mut MutateContext<'_, State>`", + )); + }; + let Some(segment) = path.path.segments.last() else { + return Err(syn::Error::new_spanned( + context_type, + "expected `&mut MutateContext<'_, State>`", + )); + }; + if segment.ident != "MutateContext" { + return Err(syn::Error::new_spanned( + context_type, + "expected `&mut MutateContext<'_, State>`", + )); + } + let PathArguments::AngleBracketed(arguments) = &segment.arguments else { + return Err(syn::Error::new_spanned( + context_type, + "`MutateContext` requires its lifetime and state type", + )); + }; + let mut state_types = arguments.args.iter().filter_map(|argument| match argument { + GenericArgument::Type(state) => Some(state.clone()), + _ => None, + }); + let Some(state) = state_types.next() else { + return Err(syn::Error::new_spanned( + context_type, + "`MutateContext` requires a state type", + )); + }; + if state_types.next().is_some() { + return Err(syn::Error::new_spanned( + context_type, + "`MutateContext` accepts exactly one state type", + )); + } + Ok(state) +} + fn presence_attrs(attrs: &[syn::Attribute]) -> syn::Result<Vec<Meta>> { attrs .iter() diff --git a/rust/tvm-ffi/src/extra/structural_mutate.rs b/rust/tvm-ffi/src/extra/structural_mutate.rs index e7088bd4..88e4a9fc 100644 --- a/rust/tvm-ffi/src/extra/structural_mutate.rs +++ b/rust/tvm-ffi/src/extra/structural_mutate.rs @@ -139,6 +139,12 @@ impl<State> MutateContext<'_, State> { self.def_region_kind } + /// Definition region active at the callback's current value. + #[inline] + pub fn region(&self) -> DefRegionKind { + self.def_region_kind + } + /// Mutate a borrowed value through the same callback chain. The value and /// its descendants begin on the non-in-place path. pub fn mutate<T>(&mut self, value: &T) -> Result<Any> @@ -200,11 +206,12 @@ impl<State> MutateContext<'_, State> { /// Conversion into the mutator argument accepted by [`structural_mutate`]. /// -/// Accepts a mutable [`StructuralMutator`] or a first-match callback chain. -/// Use [`MutateCallbacks`] when the chain needs mutable state. +/// Accepts a mutable low-level [`StructuralMutator`], a generated +/// [`MutateDispatch`], or a first-match callback chain. Use [`MutateCallbacks`] +/// when typed dispatch or a callback chain needs mutable state. #[diagnostic::on_unimplemented( message = "`{Self}` is not a supported `structural_mutate` mutator", - note = "accepted mutators: `&mut U` where `U: StructuralMutator`; an `Fn` callback over an FFI value type `T`, `&N` of an object node type, or `&MapValue`, followed by `&mut MutateContext<'_, ()>`; or a tuple of up to 12 such callbacks (tuples may nest)", + note = "accepted mutators: `&mut U` where `U: StructuralMutator`; a generated `MutateDispatch<State = ()>`; an `Fn` callback over an FFI value type `T`, `&N` of an object node type, or `&MapValue`, followed by `&mut MutateContext<'_, ()>`; or a tuple of up to 12 such callbacks (tuples may nest)", note = "callback arguments need explicit type annotations; use `MutateCallbacks::new(state, callbacks)` for ordinary mutable callback state" )] pub trait IntoMutator<Marker> { @@ -254,6 +261,23 @@ pub trait MutateChainLink<State, Marker>: mutate_sealed::SealedLink<State, Marke ) -> Option<MutateResult>; } +/// Ordered typed callback dispatch for [`structural_mutate`]. +/// +/// `None` means no handler matched, so structural mutation applies its default +/// behavior. A generated `#[dispatch(mutate)]` implementation tests +/// `mutate_*` methods in source order and passes the same [`MutateContext`] to +/// the first match. +pub trait MutateDispatch: Sized { + /// Mutable state shared by the dispatched callbacks. + type State; + + fn dispatch_mutate( + &self, + value: &MapValue, + mutator: &mut MutateContext<'_, Self::State>, + ) -> Option<MutateResult>; +} + mod mutate_sealed { use super::{IntoMutateResult, MapValue, MutateContext, ObjectCore}; @@ -285,6 +309,25 @@ mod mutate_sealed { O: IntoMutateResult, { } + + impl<D> SealedLink<D::State, super::ByMutateDispatch> for D where D: super::MutateDispatch {} +} + +#[doc(hidden)] +pub enum ByMutateDispatch {} + +impl<D> MutateChainLink<D::State, ByMutateDispatch> for D +where + D: MutateDispatch, +{ + #[inline] + fn try_mutate( + &self, + value: &MapValue, + mutator: &mut MutateContext<'_, D::State>, + ) -> Option<MutateResult> { + self.dispatch_mutate(value, mutator) + } } #[doc(hidden)] @@ -385,7 +428,7 @@ macro_rules! impl_mutate_chain_link { impl_callback_chain_tuple_arities!(impl_mutate_chain_link); -/// A reusable callback mutator with shared user state. +/// A reusable typed-dispatch or callback mutator with shared user state. pub struct MutateCallbacks<State, Link, Marker> { state: State, callbacks: Rc<Link>, @@ -888,10 +931,11 @@ impl StructuralVarRemap { } } -/// A mutator that controls its own recursion. +/// A low-level mutator that controls its own recursion. /// /// Implementations descend with the `mutate` or `default_*` helpers. -/// `#[dispatch(mutate)]` generates this trait from typed `mutate_*` methods. +/// Prefer mutation callbacks or `#[dispatch(mutate)]` for typed dispatch with +/// recursion supplied through [`MutateContext`]. pub trait StructuralMutator: Sized { /// Dispatch one borrowed value without modifying its source storage. /// diff --git a/rust/tvm-ffi/src/lib.rs b/rust/tvm-ffi/src/lib.rs index 8be52bbe..bcd19870 100644 --- a/rust/tvm-ffi/src/lib.rs +++ b/rust/tvm-ffi/src/lib.rs @@ -51,7 +51,7 @@ pub use crate::extra::module::Module; pub use crate::extra::structural_mutate::{ structural_map, structural_mutate, InplaceValue, IntoMapResult, IntoMapper, IntoMutator, MapChainLink, MapDispatch, MapValue, MutateCallbacks, MutateChainLink, MutateContext, - StructuralMutator, StructuralVarRemap, + MutateDispatch, StructuralMutator, StructuralVarRemap, }; pub use crate::extra::structural_visit::{ structural_visit, structural_walk, DefRegionKind, IntoVisitor, IntoWalkResult, IntoWalker, diff --git a/rust/tvm-ffi/tests/test_structural_mutate.rs b/rust/tvm-ffi/tests/test_structural_mutate.rs index cf364a56..c11fb2bb 100644 --- a/rust/tvm-ffi/tests/test_structural_mutate.rs +++ b/rust/tvm-ffi/tests/test_structural_mutate.rs @@ -1377,23 +1377,49 @@ fn generated_map_dispatch_supports_kind_and_ordered_catch_all() { } #[derive(Default)] -struct GeneratedLeafMutator { +struct GeneratedLeafState { integers: Vec<(i64, DefRegionKind)>, } +struct GeneratedLeafDispatch; + #[dispatch(mutate)] -impl GeneratedLeafMutator { - fn mutate_integer(&mut self, value: i64, kind: DefRegionKind) -> Any { - self.integers.push((value, kind)); +impl GeneratedLeafDispatch { + fn mutate_integer( + &self, + value: i64, + mutator: &mut MutateContext<'_, GeneratedLeafState>, + ) -> Any { + let region = mutator.region(); + mutator.state_mut().integers.push((value, region)); Any::from(value + 1) } } +struct GeneratedStatelessDispatch; + +#[dispatch(mutate)] +impl GeneratedStatelessDispatch { + fn mutate_integer(&self, value: i64, _mutator: &mut MutateContext<'_, ()>) -> i64 { + value + 1 + } +} + #[test] -fn generated_mutator_defaults_unmatched_values_and_preserves_inplace_permit() { +fn generated_stateless_mutate_dispatch_is_a_direct_callback() { + assert_eq!( + structural_mutate(1i64, GeneratedStatelessDispatch) + .and_then(i64::try_from) + .unwrap(), + 2 + ); +} + +#[test] +fn generated_mutate_dispatch_defaults_unmatched_values_and_preserves_inplace_permit() { let root = Array::new(vec![1i64, 2]); let root_pointer = array_pointer(&root); - let mut mutator = GeneratedLeafMutator::default(); + let mut mutator = MutateCallbacks::new(GeneratedLeafState::default(), GeneratedLeafDispatch); let mutated = structural_mutate(root, &mut mutator) .and_then(Array::<i64>::try_from) .unwrap(); @@ -1401,126 +1427,196 @@ fn generated_mutator_defaults_unmatched_values_and_preserves_inplace_permit() { assert_eq!(array_pointer(&mutated), root_pointer); assert_eq!(mutated.iter().collect::<Vec<_>>(), vec![2, 3]); assert_eq!( - mutator.integers, + mutator.state().integers, vec![(1, DefRegionKind::None), (2, DefRegionKind::None)] ); } #[test] -fn generated_mutator_default_remap_crosses_registered_hooks() { +fn generated_mutate_dispatch_default_remap_crosses_registered_hooks() { ensure_test_types_registered(); let _guard = REGISTERED_HOOK_TEST_LOCK.lock().unwrap(); RETAINED_MUTATOR.with(|retained| { retained.take(); }); - let mut mutator = GeneratedLeafMutator::default(); + let mut mutator = MutateCallbacks::new(GeneratedLeafState::default(), GeneratedLeafDispatch); let mutated = structural_mutate(rust_hook_node(), &mut mutator) .and_then(i64::try_from) .unwrap(); assert_eq!(mutated, 2); - assert_eq!(mutator.integers, vec![(1, DefRegionKind::None)]); + assert_eq!(mutator.state().integers, vec![(1, DefRegionKind::None)]); RETAINED_MUTATOR.with(|retained| { retained.take(); }); } #[derive(Default)] -struct GeneratedRecursiveMutator { - integers: Vec<i64>, +struct GeneratedRecursiveState { + arrays: Vec<DefRegionKind>, + integers: Vec<(i64, DefRegionKind)>, } +struct GeneratedRecursiveDispatch; + #[dispatch(mutate)] -impl GeneratedRecursiveMutator { - fn mutate_array(&mut self, array: Array<i64>, kind: DefRegionKind) -> Result<Array<i64>> { +impl GeneratedRecursiveDispatch { + fn mutate_array( + &self, + array: Array<i64>, + mutator: &mut MutateContext<'_, GeneratedRecursiveState>, + ) -> Result<Array<i64>> { + let region = mutator.region(); + mutator.state_mut().arrays.push(region); let mut mutated = Vec::with_capacity(array.len()); for value in array.iter() { - mutated.push(i64::try_from(self.mutate(&value, kind)?)?); + mutated.push(i64::try_from(mutator.mutate(&value)?)?); } Ok(Array::new(mutated)) } - fn mutate_integer(&mut self, value: i64) -> Any { - self.integers.push(value); + fn mutate_integer( + &self, + value: i64, + mutator: &mut MutateContext<'_, GeneratedRecursiveState>, + ) -> Any { + let region = mutator.region(); + mutator.state_mut().integers.push((value, region)); Any::from(value + 10) } } #[test] -fn generated_mutator_can_drive_recursion_through_mut_self() { - let mut mutator = GeneratedRecursiveMutator::default(); +fn generated_mutate_dispatch_recurses_through_context() { + let mut mutator = MutateCallbacks::new( + GeneratedRecursiveState::default(), + GeneratedRecursiveDispatch, + ); let mutated = structural_mutate(Array::new(vec![1i64, 2]), &mut mutator) .and_then(Array::<i64>::try_from) .unwrap(); assert_eq!(mutated.iter().collect::<Vec<_>>(), vec![11, 12]); - assert_eq!(mutator.integers, vec![1, 2]); + assert_eq!(mutator.state().arrays, vec![DefRegionKind::None]); + assert_eq!( + mutator.state().integers, + vec![(1, DefRegionKind::None), (2, DefRegionKind::None)] + ); +} + +#[test] +fn generated_mutate_dispatch_inherits_region_during_explicit_recursion() { + ensure_test_types_registered(); + let _guard = REFLECTED_TEST_LOCK.lock().unwrap(); + let root = rust_pair(Array::new(vec![1i64]), Any::new()); + let mut mutator = MutateCallbacks::new( + GeneratedRecursiveState::default(), + GeneratedRecursiveDispatch, + ); + + let mutated = structural_mutate(root, &mut mutator) + .and_then(RustPair::try_from) + .unwrap(); + let first = Array::<i64>::try_from(mutated.data.first.clone()).unwrap(); + + assert_eq!(first.iter().collect::<Vec<_>>(), vec![11]); + assert_eq!(mutator.state().arrays, vec![DefRegionKind::Recursive]); + assert_eq!( + mutator.state().integers, + vec![(1, DefRegionKind::Recursive)] + ); } #[derive(Default)] -struct GeneratedDefaultingMutator { +struct GeneratedDefaultingState { arrays: usize, integers: Vec<i64>, } +struct GeneratedDefaultingDispatch; + #[dispatch(mutate)] -impl GeneratedDefaultingMutator { - fn mutate_array(&mut self, array: Array<i64>, kind: DefRegionKind) -> Result<Any> { - self.arrays += 1; - self.default_mutate_value(&array, kind) +impl GeneratedDefaultingDispatch { + fn mutate_array( + &self, + _array: Array<i64>, + mutator: &mut MutateContext<'_, GeneratedDefaultingState>, + ) -> Result<Any> { + mutator.state_mut().arrays += 1; + mutator.default_mutate() } - fn mutate_integer(&mut self, value: i64) -> Any { - self.integers.push(value); + fn mutate_integer( + &self, + value: i64, + mutator: &mut MutateContext<'_, GeneratedDefaultingState>, + ) -> Any { + mutator.state_mut().integers.push(value); Any::from(value + 1) } } #[test] -fn generated_mutator_can_default_recurse_from_a_typed_handler() { - let mut mutator = GeneratedDefaultingMutator::default(); +fn generated_mutate_dispatch_can_default_recurse_from_a_typed_handler() { + let mut mutator = MutateCallbacks::new( + GeneratedDefaultingState::default(), + GeneratedDefaultingDispatch, + ); let mutated = structural_mutate(Array::new(vec![1i64, 2]), &mut mutator) .and_then(Array::<i64>::try_from) .unwrap(); assert_eq!(mutated.iter().collect::<Vec<_>>(), vec![2, 3]); - assert_eq!(mutator.arrays, 1); - assert_eq!(mutator.integers, vec![1, 2]); + assert_eq!(mutator.state().arrays, 1); + assert_eq!(mutator.state().integers, vec![1, 2]); } -struct GeneratedRemappingMutator { +struct GeneratedRemappingState { type_index: i32, calls: usize, } +struct GeneratedRemappingDispatch; + #[dispatch(mutate)] -impl GeneratedRemappingMutator { - fn mutate_dag_node(&mut self, _value: &RustDagNodeObj) -> Any { +impl GeneratedRemappingDispatch { + fn mutate_dag_node( + &self, + _value: &RustDagNodeObj, + _mutator: &mut MutateContext<'_, GeneratedRemappingState>, + ) -> Any { Any::from(42i64) } - fn mutate_any(&mut self, value: &MapValue, kind: DefRegionKind) -> Result<Any> { - if value.type_index() != self.type_index { - return self.default_mutate(value, kind); + fn mutate_any( + &self, + value: &MapValue, + mutator: &mut MutateContext<'_, GeneratedRemappingState>, + ) -> Result<Any> { + if value.type_index() != mutator.state().type_index { + return mutator.default_mutate(); } - if let Some(mutated) = self.var_remap_get(value)? { + if let Some(mutated) = mutator.var_remap_get(value)? { return Ok(mutated); } - self.calls += 1; + mutator.state_mut().calls += 1; let mutated = Any::from(41i64); - self.var_remap_set(value, &mutated)?; + mutator.var_remap_set(value, &mutated)?; Ok(mutated) } } #[test] -fn generated_mutator_uses_fresh_invocation_local_var_remap() { +fn generated_mutate_dispatch_uses_fresh_invocation_local_var_remap() { ensure_test_types_registered(); let var = rust_free_var(); - let mut mutator = GeneratedRemappingMutator { - type_index: RustFreeVarObj::type_index(), - calls: 0, - }; + let mut mutator = MutateCallbacks::new( + GeneratedRemappingState { + type_index: RustFreeVarObj::type_index(), + calls: 0, + }, + GeneratedRemappingDispatch, + ); for expected_calls in [1, 2] { let root = call_global( @@ -1528,7 +1624,7 @@ fn generated_mutator_uses_fresh_invocation_local_var_remap() { &[Any::from(var.clone()), Any::from(var.clone())], ); let mutated = structural_mutate(root, &mut mutator).unwrap(); - assert_eq!(mutator.calls, expected_calls); + assert_eq!(mutator.state().calls, expected_calls); assert_eq!(i64::try_from(array_item(&mutated, 0)).unwrap(), 41); assert_eq!(i64::try_from(array_item(&mutated, 1)).unwrap(), 41); }
