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);
     }

Reply via email to