diff --git a/docs/guides/rust_lang_guide.md b/docs/guides/rust_lang_guide.md index f4c19e786..f2ae5eff5 100644 --- a/docs/guides/rust_lang_guide.md +++ b/docs/guides/rust_lang_guide.md @@ -180,6 +180,154 @@ fn may_fail(value: i32) -> Result<()> { } ``` +### Structural Walk and Visit + +Rust provides equivalents of the C++ `StructuralWalk`/`StructuralVisitor` +APIs. Put `#[dispatch(visit)]` on an impl to turn its `visit_*` methods into +typed handlers, then pass it to `structural_walk`; each handler returns a +`WalkResult` (`Advance`, `Skip`, or `Interrupt`) to steer the traversal. +Handlers dispatch on their argument type and may take an optional trailing +`DefRegionKind` argument: + +```rust +use tvm_ffi::{dispatch, structural_walk, Array, DefRegionKind, WalkOrder, WalkResult}; + +#[derive(Default)] +struct Probe { + total: i64, + floats: usize, +} + +#[dispatch(visit)] +impl Probe { + fn visit_integer(&mut self, value: i64) -> WalkResult { + self.total += value; + WalkResult::Advance + } + + fn visit_float(&mut self, _value: f64, _kind: DefRegionKind) -> WalkResult { + self.floats += 1; + WalkResult::Advance + } +} + +let values = Array::new(vec![1_i64, 2, 3]); +let mut probe = Probe::default(); +structural_walk(&values, &mut probe, WalkOrder::PreOrder)?; +assert_eq!(probe.total, 6); +``` + +Lambdas also work — pass a single typed lambda, or a tuple of them (up to 8) +tried in order with the first matching argument type winning, like the +variadic C++ `StructuralWalk(root, callbacks...)` chain. Unmatched values +simply advance; a `&VisitValue` lambda acts as a catch-all and must come +last, since links after an always-matching one never run. Each lambda may +take a trailing `DefRegionKind` argument: + +```rust +use tvm_ffi::{structural_walk, Array, DefRegionKind, Object, WalkOrder, WalkResult}; + +let values = Array::new(vec![1_i64, 2, 3]); + +let mut total = 0; +structural_walk( + &values, + |value: i64| { + total += value; + WalkResult::Advance + }, + WalkOrder::PreOrder, +)?; +assert_eq!(total, 6); + +let mut evens = 0; +let mut objects = 0; +structural_walk( + &values, + ( + |value: i64| { + if value % 2 == 0 { + evens += 1; + } + WalkResult::Advance + }, + |_object: &Object, _kind: DefRegionKind| { + objects += 1; + WalkResult::Advance + }, + ), + WalkOrder::PreOrder, +)?; +assert_eq!((evens, objects), (1, 1)); +``` + +Both entry points return `Result>`: `Ok(None)` means +the whole graph was visited, and a handler stops the walk early by returning +`WalkResult::interrupt_with(payload)`, which comes back to the caller as +`Ok(Some(interrupt))`. Handlers may also return `Result` and +propagate errors with `?`: + +```rust +use tvm_ffi::{structural_walk, Array, WalkOrder, WalkResult}; + +let values = Array::new(vec![1_i64, 2, 3]); +let found = structural_walk( + &values, + |value: i64| { + if value == 2 { + return WalkResult::interrupt_with(value); + } + WalkResult::Advance + }, + WalkOrder::PreOrder, +)?; +assert_eq!(found.map(|i| i64::try_from(i.value).unwrap()), Some(2)); +``` + +To drive recursion yourself, implement `StructuralVisitor` and call +`structural_visit`; `visit` runs for each value and descends through +`default_visit_children`, or through `visit_child`, which visits one +selected child and can override the def-region state for it (e.g. +`DefRegionKind::Recursive` when descending into a binder's parameters): + +```rust +use tvm_ffi::{ + structural_visit, Array, DefRegionKind, Result, StructuralVisitor, VisitInterrupt, VisitValue, +}; + +#[derive(Default)] +struct Depth { + max: usize, + current: usize, +} + +impl StructuralVisitor for Depth { + fn visit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result> { + self.current += 1; + self.max = self.max.max(self.current); + let interrupt = self.default_visit_children(value, def_region_kind)?; + self.current -= 1; + Ok(interrupt) + } +} + +let values = Array::new(vec![1_i64, 2]); +let mut depth = Depth::default(); +structural_visit(&values, &mut depth)?; +assert_eq!(depth.max, 2); +``` + +Two safety notes: mutable `List`/`Dict` contents are snapshotted before +callbacks run, so mutation during traversal cannot invalidate the walk; and +a non-container type with a foreign `__s_visit__` hook is rejected rather +than silently walked through reflection — visit such a type's children +explicitly from a `StructuralVisitor`, or skip it with a pre-order +`WalkResult::Skip`. + ## Examples The repository includes a complete example in `rust/tvm-ffi/examples/load_library.rs`. diff --git a/rust/tvm-ffi-macros/src/lib.rs b/rust/tvm-ffi-macros/src/lib.rs index 03ecb050e..ada75d419 100644 --- a/rust/tvm-ffi-macros/src/lib.rs +++ b/rust/tvm-ffi-macros/src/lib.rs @@ -23,6 +23,15 @@ use proc_macro_error::proc_macro_error; mod match_any; mod object_macros; mod utils; +mod visit; + +/// Generate typed structural-visit dispatch from the `visit_*` methods in an +/// inherent implementation. +#[proc_macro_error] +#[proc_macro_attribute] +pub fn dispatch(attr: TokenStream, item: TokenStream) -> TokenStream { + visit::dispatch(attr, item) +} /// Match object-backed values carried by an Any-compatible scrutinee. /// diff --git a/rust/tvm-ffi-macros/src/visit.rs b/rust/tvm-ffi-macros/src/visit.rs new file mode 100644 index 000000000..4bad37db0 --- /dev/null +++ b/rust/tvm-ffi-macros/src/visit.rs @@ -0,0 +1,279 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +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 crate::utils::get_tvm_ffi_crate; + +pub(crate) fn dispatch(attr: TokenStream, item: TokenStream) -> TokenStream { + let _args = parse_macro_input!(attr as DispatchArgs); + let item_impl = parse_macro_input!(item as ItemImpl); + + match expand(&item_impl) { + Ok(generated) => quote!(#item_impl #generated).into(), + Err(error) => { + let error = error.to_compile_error(); + quote!(#item_impl #error).into() + } + } +} + +struct DispatchArgs; + +impl syn::parse::Parse for DispatchArgs { + fn parse(input: syn::parse::ParseStream) -> syn::Result { + let mode: syn::Ident = input.parse()?; + if mode != "visit" { + return Err(syn::Error::new(mode.span(), "expected `dispatch(visit)`")); + } + if !input.is_empty() { + return Err(input.error( + "`dispatch(visit)` takes no further arguments; a handler that needs the \ + definition-region state declares a trailing `DefRegionKind` argument", + )); + } + Ok(DispatchArgs) + } +} + +fn expand(item_impl: &ItemImpl) -> syn::Result { + if item_impl.trait_.is_some() { + return Err(syn::Error::new_spanned( + item_impl, + "`dispatch(visit)` requires an inherent impl", + )); + } + + let handlers = item_impl + .items + .iter() + .filter_map(|item| match item { + ImplItem::Method(method) if method.sig.ident.to_string().starts_with("visit_") => { + Some(parse_handler(method)) + } + _ => None, + }) + .collect::>>()?; + + if handlers.is_empty() { + return Err(syn::Error::new_spanned( + item_impl, + "`dispatch(visit)` found no `visit_*` methods", + )); + } + let tvm_ffi = get_tvm_ffi_crate(); + + let links = handlers.iter().map(|handler| { + let method = &handler.method; + let attrs = &handler.cfg_attrs; + // A handler opts into the definition-region state by declaring a + // trailing argument; the generated dispatch forwards by arity, like + // the C++ StructuralWalk callback overloads. + let kind_arg = if handler.wants_def_region { + quote!(, def_region_kind) + } else { + quote!() + }; + let invoke = match &handler.argument { + HandlerArgument::Value => quote! { + return Some( + #tvm_ffi::extra::structural_visit::IntoVisitResult::into_visit_result( + self.#method(value #kind_arg) + ) + ); + }, + HandlerArgument::BorrowedNode(node_type) => quote! { + if let Some(node) = value.as_node::<#node_type>() { + return Some( + #tvm_ffi::extra::structural_visit::IntoVisitResult::into_visit_result( + self.#method(node #kind_arg) + ) + ); + } + }, + HandlerArgument::Owned(value_type) => quote! { + if let Some(node) = value.cast::<#value_type>() { + return Some( + #tvm_ffi::extra::structural_visit::IntoVisitResult::into_visit_result( + self.#method(node #kind_arg) + ) + ); + } + }, + }; + quote! { + #(#[#attrs])* + { + #invoke + } + } + }); + let self_type = &item_impl.self_ty; + let (impl_generics, _, where_clause) = item_impl.generics.split_for_impl(); + let impl_cfg_attrs = presence_attrs(&item_impl.attrs)?; + let ordering_errors = handlers + .iter() + .enumerate() + .filter(|(_, handler)| matches!(&handler.argument, HandlerArgument::Value)) + .flat_map(|(index, handler)| { + handlers[index + 1..].iter().map(|later| { + let span = handler.method.span(); + let handler_attrs = &handler.cfg_attrs; + let later_attrs = &later.cfg_attrs; + quote_spanned! {span=> + #(#[#impl_cfg_attrs])* + #(#[#handler_attrs])* + #(#[#later_attrs])* + compile_error!( + "the `&VisitValue` catch-all handler must be last among enabled handlers" + ); + } + }) + }); + + Ok(quote! { + #(#ordering_errors)* + + #(#[#impl_cfg_attrs])* + impl #impl_generics #tvm_ffi::extra::structural_visit::VisitDispatch + for #self_type #where_clause + { + #[allow(unreachable_code, unused_variables)] + fn dispatch_visit( + &mut self, + value: &#tvm_ffi::extra::structural_visit::VisitValue, + def_region_kind: #tvm_ffi::extra::structural_visit::DefRegionKind, + ) -> Option<#tvm_ffi::extra::structural_visit::VisitResult> { + #(#links)* + None + } + } + }) +} + +struct Handler { + method: syn::Ident, + argument: HandlerArgument, + /// The handler declared a trailing `DefRegionKind` argument. + wants_def_region: bool, + cfg_attrs: Vec, +} + +enum HandlerArgument { + Value, + BorrowedNode(Type), + Owned(Type), +} + +fn parse_handler(method: &ImplItemMethod) -> syn::Result { + 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, + "visit handlers must take `&mut self`, a node, and optionally a trailing \ + `DefRegionKind` argument", + )); + } + let wants_def_region = inputs.len() == 3; + + let value_type = match inputs.iter().nth(1) { + Some(FnArg::Typed(value)) => (*value.ty).clone(), + _ => unreachable!("the second argument cannot be a receiver"), + }; + let argument = match &value_type { + Type::Reference(reference) if reference.mutability.is_none() => { + if is_visit_value(reference.elem.as_ref()) { + HandlerArgument::Value + } else { + HandlerArgument::BorrowedNode((*reference.elem).clone()) + } + } + Type::Reference(_) => { + return Err(syn::Error::new_spanned( + value_type, + "visit handler values cannot be mutable references", + )); + } + _ => HandlerArgument::Owned(value_type), + }; + let cfg_attrs = presence_attrs(&method.attrs)?; + Ok(Handler { + method: method.sig.ident.clone(), + argument, + wants_def_region, + cfg_attrs, + }) +} + +fn presence_attrs(attrs: &[syn::Attribute]) -> syn::Result> { + attrs + .iter() + .filter(|attr| attr.path.is_ident("cfg") || attr.path.is_ident("cfg_attr")) + .map(|attr| attr.parse_meta().map(presence_meta)) + .filter_map(Result::transpose) + .collect() +} + +fn presence_meta(meta: Meta) -> Option { + if meta.path().is_ident("cfg") { + return Some(meta); + } + let Meta::List(mut list) = meta else { + return None; + }; + if !list.path.is_ident("cfg_attr") { + return None; + } + + let mut items = list.nested.into_iter(); + let condition = items.next()?; + let mut retained = syn::punctuated::Punctuated::new(); + retained.push(condition); + for item in items { + if let NestedMeta::Meta(meta) = item { + if let Some(meta) = presence_meta(meta) { + retained.push(NestedMeta::Meta(meta)); + } + } + } + if retained.len() == 1 { + None + } else { + list.nested = retained; + Some(Meta::List(list)) + } +} + +fn is_visit_value(value_type: &Type) -> bool { + let Type::Path(path) = value_type else { + return false; + }; + path.path + .segments + .last() + .is_some_and(|segment| segment.ident == "VisitValue") +} diff --git a/rust/tvm-ffi-sys/src/c_api.rs b/rust/tvm-ffi-sys/src/c_api.rs index 910a0505a..203446c89 100644 --- a/rust/tvm-ffi-sys/src/c_api.rs +++ b/rust/tvm-ffi-sys/src/c_api.rs @@ -81,6 +81,16 @@ pub enum TVMFFITypeIndex { kTVMFFIModule = 73, /// Opaque python object. kTVMFFIOpaquePyObject = 74, + /// Mutable list object. + kTVMFFIList = 75, + /// Mutable dict object. + kTVMFFIDict = 76, + /// Structural visit interrupt object. + kTVMFFIVisitInterrupt = 77, + /// End of the statically allocated object type-index range. + kTVMFFIStaticObjectEnd = 78, + /// Start of dynamically allocated object type indices. + kTVMFFIDynObjectBegin = 128, } #[repr(i32)] @@ -91,6 +101,46 @@ pub enum TVMFFIObjectDeleterFlagBitMask { kTVMFFIObjectDeleterFlagBitMaskBoth = (1 << 0) | (1 << 1), } +/// Bit flags attached to reflected fields. +#[repr(i32)] +#[derive(Debug, Copy, Clone, PartialEq, Eq)] +pub enum TVMFFIFieldFlagBitMask { + kTVMFFIFieldFlagBitMaskWritable = 1 << 0, + kTVMFFIFieldFlagBitMaskHasDefault = 1 << 1, + kTVMFFIFieldFlagBitMaskIsStaticMethod = 1 << 2, + kTVMFFIFieldFlagBitMaskSEqHashIgnore = 1 << 3, + kTVMFFIFieldFlagBitMaskSEqHashDefRecursive = 1 << 4, + kTVMFFIFieldFlagBitMaskDefaultFromFactory = 1 << 5, + kTVMFFIFieldFlagBitMaskReprOff = 1 << 6, + kTVMFFIFieldFlagBitMaskCompareOff = 1 << 7, + kTVMFFIFieldFlagBitMaskHashOff = 1 << 8, + kTVMFFIFieldFlagBitMaskInitOff = 1 << 9, + kTVMFFIFieldFlagBitMaskKwOnly = 1 << 10, + kTVMFFIFieldFlagBitSetterIsFunctionObj = 1 << 11, + kTVMFFIFieldFlagBitMaskSEqHashDefNonRecursive = 1 << 12, +} + +/// Definition-region mode used by structural traversal. +#[repr(i32)] +#[derive(Debug, Copy, Clone, PartialEq, Eq)] +pub enum TVMFFIDefRegionKind { + kTVMFFIDefRegionKindNone = 0, + kTVMFFIDefRegionKindRecursive = 1, + kTVMFFIDefRegionKindNonRecursive = 2, +} + +/// Structural equality/hash participation kind stored in type metadata. +#[repr(i32)] +#[derive(Debug, Copy, Clone, PartialEq, Eq)] +pub enum TVMFFISEqHashKind { + kTVMFFISEqHashKindUnsupported = 0, + kTVMFFISEqHashKindTreeNode = 1, + kTVMFFISEqHashKindFreeVar = 2, + kTVMFFISEqHashKindDAGNode = 3, + kTVMFFISEqHashKindConstTreeNode = 4, + kTVMFFISEqHashKindUniqueInstance = 5, +} + /// Handle to Object from C API's pov pub type TVMFFIObjectHandle = *mut c_void; pub type TVMFFIObjectDeleter = unsafe extern "C" fn(self_ptr: *mut c_void, flags: i32); @@ -410,6 +460,14 @@ unsafe extern "C" { pub fn TVMFFISetCustomAllocator(allocator: *mut TVMFFICustomAllocator) -> i32; pub fn TVMFFITypeKeyToIndex(type_key: *const TVMFFIByteArray, out_tindex: *mut i32) -> i32; + pub fn TVMFFITypeRegisterAttr( + type_index: i32, + attr_name: *const TVMFFIByteArray, + attr_value: *const TVMFFIAny, + ) -> i32; + pub fn TVMFFIGetTypeAttrColumn( + attr_name: *const TVMFFIByteArray, + ) -> *const TVMFFITypeAttrColumn; pub fn TVMFFIFunctionGetGlobal( name: *const TVMFFIByteArray, out: *mut TVMFFIObjectHandle, diff --git a/rust/tvm-ffi/src/any.rs b/rust/tvm-ffi/src/any.rs index 47379f47a..1ebf017c7 100644 --- a/rust/tvm-ffi/src/any.rs +++ b/rust/tvm-ffi/src/any.rs @@ -53,6 +53,25 @@ impl<'a> AnyView<'a> { self.data.type_index } + #[inline] + pub(crate) fn as_raw_ffi_any(&self) -> &TVMFFIAny { + &self.data + } + + /// Construct a borrowed view from its ABI representation. + /// + /// # Safety + /// + /// The caller must keep every resource referenced by `data` alive for the + /// returned view's complete lifetime. + #[inline] + pub(crate) unsafe fn from_raw_ffi_any(data: TVMFFIAny) -> Self { + Self { + data, + _phantom: std::marker::PhantomData, + } + } + /// More strict version than try_from/try_into /// /// This function will not try to cast the type @@ -151,6 +170,10 @@ impl Any { pub fn type_index(&self) -> i32 { self.data.type_index } + #[inline] + pub(crate) fn as_raw_ffi_any(&self) -> &TVMFFIAny { + &self.data + } /// Try to query if stored typed in Any exactly matches the type T /// /// This function is fast in the case of failure and can be used to check diff --git a/rust/tvm-ffi/src/extra/dispatch.rs b/rust/tvm-ffi/src/extra/dispatch.rs new file mode 100644 index 000000000..2b3a49f0b --- /dev/null +++ b/rust/tvm-ffi/src/extra/dispatch.rs @@ -0,0 +1,114 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +//! Typed visitor dispatch for [`super::structural_visit::structural_walk`]: +//! the [`VisitDispatch`] trait targeted by `#[dispatch(visit)]`, and the +//! walker adapter that runs such a visitor at the phase selected by the walk +//! order. The traversal engine, closure walkers, and tuple chains live in +//! [`super::structural_visit`], which re-exports these items to keep its +//! public paths stable. + +use crate::error::Result; + +use super::structural_visit::{ + DefRegionKind, IntoWalker, NativeVisit, VisitResult, VisitValue, WalkOrder, WalkResult, +}; + +/// Typed dispatch implemented by a walk-layer observer. +/// +/// [`crate::dispatch`] tests the implementation's `visit_*` methods in source +/// order. Borrowed node arguments use refcount-free subtype checks, owned +/// FFI-compatible arguments use exact value casts, and `&VisitValue` is a +/// catch-all. `None` reports that no handler matched: a standalone walk then +/// advances normally, while a tuple chain hands the value to the next link — +/// so a spliced visitor that "handled" a value must not return `None`. +/// +/// This is the observer layer, mirroring C++ `StructuralWalk` callbacks: the +/// walker owns recursion, and a handler steers it only through the returned +/// [`WalkResult`]. A traversal that must visit children itself — selected +/// children, custom orders, explicit definition-region overrides — belongs in +/// a [`super::structural_visit::StructuralVisitor`] instead. +/// +/// The definition-region state active at the dispatched value arrives as the +/// `def_region_kind` argument. A `#[dispatch(visit)]` handler opts into it by +/// declaring a trailing `DefRegionKind` parameter — the analog of a C++ +/// `StructuralWalk` callback accepting `(value, def_region_kind)` instead of +/// `(value)`. +pub trait VisitDispatch: Sized { + fn dispatch_visit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Option; +} + +impl VisitDispatch for &mut V { + #[inline] + fn dispatch_visit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Option { + (**self).dispatch_visit(value, def_region_kind) + } +} + +#[doc(hidden)] +pub enum ByDispatch {} + +impl<'a, V: VisitDispatch> IntoWalker for &'a mut V { + type Walker = DispatchVisitor<&'a mut V>; + fn into_walker(self, order: WalkOrder) -> Self::Walker { + DispatchVisitor { + visitor: self, + order, + } + } +} + +/// Owns its walker so a closure's state stays inline and a `&mut` visitor +/// keeps a single level of indirection. Public only as an +/// [`IntoWalker::Walker`] projection. +#[doc(hidden)] +pub struct DispatchVisitor { + visitor: V, + order: WalkOrder, +} + +impl NativeVisit for DispatchVisitor { + fn enter(&mut self, value: &VisitValue, def_region_kind: DefRegionKind) -> Result { + match self.order { + WalkOrder::PreOrder => self + .visitor + .dispatch_visit(value, def_region_kind) + .unwrap_or(Ok(WalkResult::Advance)), + WalkOrder::PostOrder => Ok(WalkResult::Advance), + } + } + + fn exit(&mut self, value: &VisitValue, def_region_kind: DefRegionKind) -> Result { + match self.order { + WalkOrder::PreOrder => Ok(WalkResult::Advance), + WalkOrder::PostOrder => self + .visitor + .dispatch_visit(value, def_region_kind) + .unwrap_or(Ok(WalkResult::Advance)), + } + } +} diff --git a/rust/tvm-ffi/src/extra/mod.rs b/rust/tvm-ffi/src/extra/mod.rs index 2a4c01d0b..0489fe27f 100644 --- a/rust/tvm-ffi/src/extra/mod.rs +++ b/rust/tvm-ffi/src/extra/mod.rs @@ -16,4 +16,6 @@ * specific language governing permissions and limitations * under the License. */ +pub mod dispatch; pub mod module; +pub mod structural_visit; diff --git a/rust/tvm-ffi/src/extra/structural_visit.rs b/rust/tvm-ffi/src/extra/structural_visit.rs new file mode 100644 index 000000000..1d15efd89 --- /dev/null +++ b/rust/tvm-ffi/src/extra/structural_visit.rs @@ -0,0 +1,1750 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +//! Native Rust structural visiting. +//! +//! Two public layers mirror the C++ API split: +//! +//! * [`StructuralVisitor`] + [`structural_visit`] — the visitor drives +//! recursion itself, like a hand-written C++ `StructuralVisitorObj`: +//! [`StructuralVisitor::visit`] runs once per reached value and descends +//! only where it calls [`StructuralVisitor::default_visit_children`] or +//! [`StructuralVisitor::visit_child`]. +//! * [`VisitDispatch`] + [`structural_walk`] — observer callbacks, like C++ +//! `StructuralWalk`: the walker recurses on its own and callbacks steer it +//! through the returned [`WalkResult`] (advance, skip, interrupt). +//! +//! Both layers thread the definition-region state explicitly: walk handlers +//! opt in with a trailing [`DefRegionKind`] argument, and a visitor receives +//! and forwards it when descending. +//! +//! Underneath both, [`VisitValue`] provides borrowed matching for typed Rust +//! dispatch and the stateless recursion engine (`visit_raw` and the +//! `visit_*` helpers below) owns iteration over containers and reflected +//! fields. +//! +//! The runtime object registry is open, so the walker uses the stable tvm-ffi +//! reflection ABI for arbitrary registered object types. That ABI is only the +//! object-description boundary: traversal, control flow, typed dispatch, +//! visitor state, and definition-region propagation remain in Rust. +//! +//! Mutable `List`/`Dict` contents are snapshotted before callbacks run, so a +//! callback mutating the container it was reached through cannot invalidate +//! the traversal; the walk sees the pre-mutation contents. +//! +//! No C++ `ffi.StructuralVisitor` is constructed and no C++ default-visit +//! function is called. A non-container type with a foreign `__s_visit__` hook +//! is rejected instead of silently substituting reflection with potentially +//! different semantics; visit such a type's children explicitly from a +//! [`StructuralVisitor`], or skip the value in a walk. + +use std::marker::PhantomData; +use std::ops::ControlFlow; +use std::os::raw::c_void; +use std::ptr::NonNull; +use std::sync::atomic::{AtomicU8, AtomicUsize, Ordering}; + +use crate::any::{Any, AnyView}; +use crate::error::{Error, Result, RUNTIME_ERROR, TYPE_ERROR}; +use crate::function::Function; +use crate::object::ObjectCore; +use crate::tvm_ffi_sys::TVMFFIFieldFlagBitMask::{ + kTVMFFIFieldFlagBitMaskSEqHashDefNonRecursive, kTVMFFIFieldFlagBitMaskSEqHashDefRecursive, + kTVMFFIFieldFlagBitMaskSEqHashIgnore, +}; +use crate::tvm_ffi_sys::{ + TVMFFIAny, TVMFFIByteArray, TVMFFIDefRegionKind, TVMFFIFieldInfo, TVMFFIGetTypeAttrColumn, + TVMFFIGetTypeInfo, TVMFFIObject, TVMFFISEqHashKind, TVMFFITypeAttrColumn, TVMFFITypeIndex, +}; + +const STRUCTURAL_VISIT_ATTR: &str = "__s_visit__"; +const FLAG_SEQ_HASH_IGNORE: i64 = kTVMFFIFieldFlagBitMaskSEqHashIgnore as i64; +const FLAG_SEQ_HASH_DEF_RECURSIVE: i64 = kTVMFFIFieldFlagBitMaskSEqHashDefRecursive as i64; +const FLAG_SEQ_HASH_DEF_NON_RECURSIVE: i64 = kTVMFFIFieldFlagBitMaskSEqHashDefNonRecursive as i64; + +/// What a callback asks the Rust walker to do with the current value. +pub enum WalkResult { + /// Continue and visit this value's children. + Advance, + /// Continue without visiting this value's children or firing its exit hook. + Skip, + /// Halt the entire traversal. + Interrupt, + /// Halt the entire traversal and return a payload to the caller. + InterruptWith(Any), +} + +impl WalkResult { + /// Halt traversal with an FFI-compatible payload. + pub fn interrupt_with>(payload: T) -> Self { + Self::InterruptWith(payload.into()) + } +} + +/// Convert either an infallible or fallible typed handler result. +/// +/// This keeps simple handlers terse while allowing a handler to return +/// `tvm_ffi::Result` and use `?`. +pub trait IntoVisitResult { + fn into_visit_result(self) -> Result; +} + +impl IntoVisitResult for WalkResult { + fn into_visit_result(self) -> Result { + Ok(self) + } +} + +impl IntoVisitResult for Result { + fn into_visit_result(self) -> Result { + self + } +} + +/// Callback order for [`structural_walk`]. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum WalkOrder { + /// Run the typed handler before the current value's children. + #[default] + PreOrder, + /// Run the typed handler after the current value's children. + PostOrder, +} + +/// Definition-region state active at the current value. +/// +/// Reflected fields marked `SEqHashDefRecursive` or +/// `SEqHashDefNonRecursive` override the inherited state for that field's +/// complete recursive visit. +#[repr(i32)] +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum DefRegionKind { + /// The value is outside a definition region. + #[default] + None = 0, + /// Definitions apply recursively through the visited value. + Recursive = 1, + /// Definitions apply to the visited value using non-recursive semantics. + NonRecursive = 2, +} + +const _: () = { + assert!(DefRegionKind::None as i32 == TVMFFIDefRegionKind::kTVMFFIDefRegionKindNone as i32); + assert!( + DefRegionKind::Recursive as i32 + == TVMFFIDefRegionKind::kTVMFFIDefRegionKindRecursive as i32 + ); + assert!( + DefRegionKind::NonRecursive as i32 + == TVMFFIDefRegionKind::kTVMFFIDefRegionKindNonRecursive as i32 + ); +}; + +/// Interrupt state of a traversal, mirroring C++ `ffi.VisitInterrupt`. +/// +/// Entry points and visitor-layer calls return +/// `Result>`: `Ok(None)` means the (sub)graph was +/// visited completely, `Ok(Some(..))` means a handler halted the traversal +/// with this interrupt, and `Err` means it failed. +pub struct VisitInterrupt { + /// Payload returned with the interrupt, or FFI `None` for no payload. + pub value: Any, +} + +impl VisitInterrupt { + /// Interrupt carrying an FFI-compatible payload. + pub fn with>(payload: T) -> Self { + Self { + value: payload.into(), + } + } +} + +/// Fallible result returned by generated typed dispatch. +#[doc(hidden)] +pub type VisitResult = Result; + +/// A borrowed view of a raw tvm-ffi value. +/// +/// Generated visitors match this value without taking ownership: borrowed +/// object-node handlers use [`VisitValue::as_node`], while POD or object-ref +/// value handlers use [`VisitValue::cast`]. +#[repr(transparent)] +pub struct VisitValue(TVMFFIAny); + +impl VisitValue { + #[inline] + fn from_raw(raw: TVMFFIAny) -> Self { + VisitValue(raw) + } + + /// Convert the value into an owned typed handle. + #[inline] + pub fn cast(&self) -> Option { + unsafe { + if R::check_any_strict(&self.0) { + Some(R::copy_from_any_view_after_check(&self.0)) + } else { + None + } + } + } + + /// Runtime type index stored in this value. + #[inline] + pub fn type_index(&self) -> i32 { + self.0.type_index + } + + /// Borrow the value as node type `N` if it is an instance of that type. + #[inline] + pub fn as_node(&self) -> Option<&N> { + if self.0.type_index < TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 { + return None; + } + let base_type_index = N::type_index(); + if self.0.type_index != base_type_index { + // A final type has no registered subtype, so a differing index can + // never match: reject with the integer compare alone, mirroring the + // `_type_final` fast path of C++ `IsObjectInstance`. + if N::TYPE_FINAL { + return None; + } + if !is_instance_at_depth(self.0.type_index, base_type_index, N::TYPE_DEPTH) { + return None; + } + } + Some(unsafe { &*(self.0.data_union.v_obj as *const N) }) + } +} + +enum NativeHalt { + Interrupt(Any), + Error(Error), +} + +impl From for NativeHalt { + fn from(error: Error) -> Self { + NativeHalt::Error(error) + } +} + +type NativeResult = std::result::Result<(), NativeHalt>; + +// The typed-dispatch layer (`VisitDispatch`, its walker adapter, and the +// `&mut V` IntoWalker form) lives in `super::dispatch`; re-exported here so +// the module's public paths — which `#[dispatch(visit)]`-generated code +// names — stay stable. +pub use super::dispatch::{ByDispatch, DispatchVisitor, VisitDispatch}; + +/// Conversion into the walker argument of [`structural_walk`]. +/// +/// The `Marker` parameter lets one entry point accept several handler +/// shapes without overlapping implementations — the Rust equivalent of the +/// C++ `StructuralWalk` callback overload set: +/// +/// * `&mut V` where `V: VisitDispatch` — a stateful typed visitor +/// (`#[dispatch(visit)]` or hand-written). +/// * A bare closure in any [`WalkChainLink`] shape — catch-all +/// `FnMut(&VisitValue)`, typed `FnMut(T)`, node `FnMut(&N)`, each with an +/// optional trailing [`DefRegionKind`] argument — the analog of a single +/// C++ callback. Values a typed closure does not match advance normally. +/// * A tuple of typed links `(link1, link2, ...)`, up to 8 — the analog of +/// the C++ variadic callback chain; see [`WalkChainLink`] for the +/// accepted link shapes. Larger handler sets belong in one +/// `#[dispatch(visit)]` visitor, which itself splices into a tuple as a +/// single link. +/// +/// Closure arguments usually need explicit type annotations +/// (`|value: &VisitValue| ...`) for the marker to be inferred. +#[diagnostic::on_unimplemented( + message = "`{Self}` is not a supported `structural_walk` walker", + note = "accepted walkers: `&mut V` where `V: VisitDispatch`; a closure over `&VisitValue`, \ + an FFI value type `T`, or `&N` of an object node type (`N: ObjectCore`, e.g. \ + `&Object`), optionally with a trailing `DefRegionKind` argument; or a tuple of \ + up to 8 such links", + note = "closure arguments need explicit type annotations; ObjectRef wrappers like `String` \ + or `Array` are FFI value types — take them by value, not by reference" +)] +pub trait IntoWalker { + #[doc(hidden)] + type Walker: NativeVisit; + #[doc(hidden)] + fn into_walker(self, order: WalkOrder) -> Self::Walker; +} + +/// Runs a catch-all closure at the phase selected by `order` — the closure +/// analog of `DispatchVisitor`, without the `Option` +/// no-handler-matched layer a dispatch chain needs. (Routing closures +/// through `DispatchVisitor` instead measures ~10-20% slower on the bare +/// closure walk: the wrapped-and-unwrapped `Option>` does not +/// fold away.) +#[doc(hidden)] +pub struct ClosureWalker { + callback: F, + order: WalkOrder, +} + +impl NativeVisit for ClosureWalker +where + F: FnMut(&VisitValue) -> O, + O: IntoVisitResult, +{ + fn enter(&mut self, value: &VisitValue, _def_region_kind: DefRegionKind) -> Result { + match self.order { + WalkOrder::PreOrder => (self.callback)(value).into_visit_result(), + WalkOrder::PostOrder => Ok(WalkResult::Advance), + } + } + + fn exit(&mut self, value: &VisitValue, _def_region_kind: DefRegionKind) -> Result { + match self.order { + WalkOrder::PreOrder => Ok(WalkResult::Advance), + WalkOrder::PostOrder => (self.callback)(value).into_visit_result(), + } + } +} + +#[doc(hidden)] +pub enum ByValueClosure {} + +impl IntoWalker for F +where + F: FnMut(&VisitValue) -> O, + O: IntoVisitResult, +{ + type Walker = ClosureWalker; + fn into_walker(self, order: WalkOrder) -> Self::Walker { + ClosureWalker { + callback: self, + order, + } + } +} + +/// `ClosureWalker` variant whose callback also receives the definition-region +/// state. +#[doc(hidden)] +pub struct ClosureKindWalker { + callback: F, + order: WalkOrder, +} + +impl NativeVisit for ClosureKindWalker +where + F: FnMut(&VisitValue, DefRegionKind) -> O, + O: IntoVisitResult, +{ + fn enter(&mut self, value: &VisitValue, def_region_kind: DefRegionKind) -> Result { + match self.order { + WalkOrder::PreOrder => (self.callback)(value, def_region_kind).into_visit_result(), + WalkOrder::PostOrder => Ok(WalkResult::Advance), + } + } + + fn exit(&mut self, value: &VisitValue, def_region_kind: DefRegionKind) -> Result { + match self.order { + WalkOrder::PreOrder => Ok(WalkResult::Advance), + WalkOrder::PostOrder => (self.callback)(value, def_region_kind).into_visit_result(), + } + } +} + +#[doc(hidden)] +pub enum ByValueKindClosure {} + +impl IntoWalker for F +where + F: FnMut(&VisitValue, DefRegionKind) -> O, + O: IntoVisitResult, +{ + type Walker = ClosureKindWalker; + fn into_walker(self, order: WalkOrder) -> Self::Walker { + ClosureKindWalker { + callback: self, + order, + } + } +} + +/// One typed link of a tuple walker — a single callback of the C++ variadic +/// `StructuralWalk(root, callbacks...)` chain. +/// +/// A tuple of up to 8 links passed to [`structural_walk`] is tried in order +/// and the first link whose argument type matches the value runs, exactly +/// like the C++ callback chain. (Python's `structural_walk` differs on one +/// point: it keeps `callbacks` and `with_def_region_kind` as two separately +/// ordered groups, trying every plain entry before any kind-taking entry, +/// so a mixed Rust tuple's single interleaved order has no exact Python +/// equivalent.) Accepted link shapes mirror `#[dispatch(visit)]` handlers: +/// +/// * `FnMut(T) -> impl IntoVisitResult` for an FFI-convertible `T` — value +/// cast via [`VisitValue::cast`], which matches on the FFI type tag: a +/// numeric link claims every `Int` (or `Float`) regardless of width and +/// converts with `as` semantics, so prefer `i64`/`f64` links unless a +/// deliberate narrowing is wanted. +/// * `FnMut(&N) -> impl IntoVisitResult` for an object node `N` — +/// refcount-free subtype check via [`VisitValue::as_node`]. +/// * `FnMut(&VisitValue) -> impl IntoVisitResult` — catch-all. +/// * `&mut V` where `V: VisitDispatch` — splice a typed visitor into the +/// chain; it claims every value one of its handlers matches. +/// +/// Links after one that matches every value never run: place a catch-all +/// closure — or a spliced visitor whose own chain ends in a `&VisitValue` +/// handler — last. Unlike the in-visitor ordering check, misordering a +/// tuple is not a compile error. +/// +/// Every closure shape may declare a trailing [`DefRegionKind`] argument, +/// and a single typed closure may also be passed to [`structural_walk`] +/// bare, without the tuple. Closure arguments need explicit type +/// annotations for the marker to be inferred. Borrow rules apply per link, +/// so state shared across links goes through a `Cell`/`RefCell` — or in a +/// single `#[dispatch(visit)]` visitor, which shares `&mut self` between +/// its handlers. +/// +/// This trait is sealed: the link shapes above are the complete set, and +/// the dispatch method is an internal detail. +pub trait WalkChainLink: sealed::SealedLink { + /// Run this link if `value` matches its argument type; `None` hands the + /// value to the next link. + #[doc(hidden)] + fn try_call( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Option; +} + +mod sealed { + use super::{DefRegionKind, IntoVisitResult, ObjectCore, VisitDispatch, VisitValue}; + + /// Seal for [`super::WalkChainLink`]: one impl per accepted link shape, + /// mirroring the `WalkChainLink` impl set exactly. + pub trait SealedLink {} + + impl SealedLink> for F + where + F: FnMut(T) -> O, + O: IntoVisitResult, + { + } + impl SealedLink> for F + where + F: FnMut(T, DefRegionKind) -> O, + O: IntoVisitResult, + { + } + impl SealedLink> for F + where + F: for<'a> FnMut(&'a N) -> O, + O: IntoVisitResult, + { + } + impl SealedLink> for F + where + F: for<'a> FnMut(&'a N, DefRegionKind) -> O, + O: IntoVisitResult, + { + } + impl SealedLink for F + where + F: for<'a> FnMut(&'a VisitValue) -> O, + O: IntoVisitResult, + { + } + impl SealedLink for F + where + F: for<'a> FnMut(&'a VisitValue, DefRegionKind) -> O, + O: IntoVisitResult, + { + } + impl SealedLink for &mut V {} +} + +#[doc(hidden)] +pub struct ByOwnedLink(PhantomData); + +impl WalkChainLink> for F +where + F: FnMut(T) -> O, + T: crate::type_traits::AnyCompatible, + O: IntoVisitResult, +{ + #[inline] + fn try_call( + &mut self, + value: &VisitValue, + _def_region_kind: DefRegionKind, + ) -> Option { + value + .cast::() + .map(|typed| self(typed).into_visit_result()) + } +} + +#[doc(hidden)] +pub struct ByOwnedKindLink(PhantomData); + +impl WalkChainLink> for F +where + F: FnMut(T, DefRegionKind) -> O, + T: crate::type_traits::AnyCompatible, + O: IntoVisitResult, +{ + #[inline] + fn try_call( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Option { + value + .cast::() + .map(|typed| self(typed, def_region_kind).into_visit_result()) + } +} + +#[doc(hidden)] +pub struct ByNodeLink(PhantomData); + +impl WalkChainLink> for F +where + F: for<'a> FnMut(&'a N) -> O, + N: ObjectCore, + O: IntoVisitResult, +{ + #[inline] + fn try_call( + &mut self, + value: &VisitValue, + _def_region_kind: DefRegionKind, + ) -> Option { + value + .as_node::() + .map(|node| self(node).into_visit_result()) + } +} + +#[doc(hidden)] +pub struct ByNodeKindLink(PhantomData); + +impl WalkChainLink> for F +where + F: for<'a> FnMut(&'a N, DefRegionKind) -> O, + N: ObjectCore, + O: IntoVisitResult, +{ + #[inline] + fn try_call( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Option { + value + .as_node::() + .map(|node| self(node, def_region_kind).into_visit_result()) + } +} + +#[doc(hidden)] +pub enum ByCatchAllLink {} + +impl WalkChainLink for F +where + F: for<'a> FnMut(&'a VisitValue) -> O, + O: IntoVisitResult, +{ + #[inline] + fn try_call( + &mut self, + value: &VisitValue, + _def_region_kind: DefRegionKind, + ) -> Option { + Some(self(value).into_visit_result()) + } +} + +#[doc(hidden)] +pub enum ByCatchAllKindLink {} + +impl WalkChainLink for F +where + F: for<'a> FnMut(&'a VisitValue, DefRegionKind) -> O, + O: IntoVisitResult, +{ + #[inline] + fn try_call( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Option { + Some(self(value, def_region_kind).into_visit_result()) + } +} + +#[doc(hidden)] +pub enum ByDispatchLink {} + +impl WalkChainLink for &mut V { + #[inline] + fn try_call( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Option { + self.dispatch_visit(value, def_region_kind) + } +} + +/// Runs a tuple of [`WalkChainLink`]s at the phase selected by `order`, +/// trying links in order and short-circuiting on the first whose type +/// matches — the Rust analog of C++ `StructuralWalkCallbackChain`. Static +/// dispatch throughout: each link's type test inlines to the same code the +/// `#[dispatch(visit)]` macro generates for a `visit_*` chain. +#[doc(hidden)] +pub struct ChainWalker { + links: Links, + order: WalkOrder, + markers: PhantomData, +} + +macro_rules! impl_chain_walker { + ($(($F:ident, $M:ident, $idx:tt)),+) => { + impl<$($F, $M,)+> ChainWalker<($($F,)+), ($($M,)+)> + where + $($F: WalkChainLink<$M>,)+ + { + #[inline] + fn dispatch( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result { + $( + if let Some(result) = self.links.$idx.try_call(value, def_region_kind) { + return result; + } + )+ + Ok(WalkResult::Advance) + } + } + + impl<$($F, $M,)+> NativeVisit for ChainWalker<($($F,)+), ($($M,)+)> + where + $($F: WalkChainLink<$M>,)+ + { + fn enter( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result { + match self.order { + WalkOrder::PreOrder => self.dispatch(value, def_region_kind), + WalkOrder::PostOrder => Ok(WalkResult::Advance), + } + } + + fn exit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result { + match self.order { + WalkOrder::PreOrder => Ok(WalkResult::Advance), + WalkOrder::PostOrder => self.dispatch(value, def_region_kind), + } + } + } + + impl<$($F, $M,)+> IntoWalker<($($M,)+)> for ($($F,)+) + where + $($F: WalkChainLink<$M>,)+ + { + type Walker = ChainWalker<($($F,)+), ($($M,)+)>; + fn into_walker(self, order: WalkOrder) -> Self::Walker { + ChainWalker { + links: self, + order, + markers: PhantomData, + } + } + } + }; +} + +impl_chain_walker!((F0, M0, 0)); +impl_chain_walker!((F0, M0, 0), (F1, M1, 1)); +impl_chain_walker!((F0, M0, 0), (F1, M1, 1), (F2, M2, 2)); +impl_chain_walker!((F0, M0, 0), (F1, M1, 1), (F2, M2, 2), (F3, M3, 3)); +impl_chain_walker!( + (F0, M0, 0), + (F1, M1, 1), + (F2, M2, 2), + (F3, M3, 3), + (F4, M4, 4) +); +impl_chain_walker!( + (F0, M0, 0), + (F1, M1, 1), + (F2, M2, 2), + (F3, M3, 3), + (F4, M4, 4), + (F5, M5, 5) +); +impl_chain_walker!( + (F0, M0, 0), + (F1, M1, 1), + (F2, M2, 2), + (F3, M3, 3), + (F4, M4, 4), + (F5, M5, 5), + (F6, M6, 6) +); +impl_chain_walker!( + (F0, M0, 0), + (F1, M1, 1), + (F2, M2, 2), + (F3, M3, 3), + (F4, M4, 4), + (F5, M5, 5), + (F6, M6, 6), + (F7, M7, 7) +); + +// A bare typed closure — `FnMut(T)` or `FnMut(&N)`, optionally with a +// trailing `DefRegionKind` — walks as a single-link chain, so a lone typed +// handler needs no tuple wrapping; values that do not match its argument +// type advance normally. `&VisitValue` catch-all closures keep their +// dedicated `ClosureWalker`/`ClosureKindWalker` path above. +macro_rules! impl_bare_link_walker { + ($(($marker:ident, $($fn_args:ty),+)),+ $(,)?) => { + $( + impl IntoWalker<$marker> for F + where + F: FnMut($($fn_args),+) -> O, + Self: WalkChainLink<$marker>, + O: IntoVisitResult, + { + type Walker = ChainWalker<(F,), ($marker,)>; + fn into_walker(self, order: WalkOrder) -> Self::Walker { + ChainWalker { + links: (self,), + order, + markers: PhantomData, + } + } + } + )+ + }; +} + +impl_bare_link_walker!( + (ByOwnedLink, T), + (ByOwnedKindLink, T, DefRegionKind), + (ByNodeLink, &T), + (ByNodeKindLink, &T, DefRegionKind), +); + +/// A visitor that drives recursion itself, mirroring C++ +/// `StructuralVisitorObj`. +/// +/// [`structural_visit`] calls [`StructuralVisitor::visit`] for the root; +/// after that the visitor is in control, exactly like a C++ visitor whose +/// vtable `visit` runs per value. A `visit` implementation descends only +/// where it chooses: +/// +/// * [`StructuralVisitor::default_visit_children`] delegates the default +/// child recursion — the analog of C++ +/// `StructuralVisitorObj::DefaultVisitExpected`. +/// * [`StructuralVisitor::visit_child`] visits one selected child — the +/// analog of C++ `visitor->Visit(child)`, with the explicit +/// `def_region_kind` argument playing the role of `WithDefRegionKind`. +/// +/// Returning without descending skips the value's children. There is no +/// [`WalkResult`] at this layer: control flow is what the implementation +/// visits, and `Ok(Some(interrupt))` halts the traversal — the analog of +/// returning a C++ `VisitInterrupt`. Nested `visit_child` and +/// `default_visit_children` calls report a nested interrupt through their +/// return value; propagate it (and errors, via `?`) upward instead of +/// dropping the result. +/// +/// The definition-region state is threaded explicitly, exactly like walk +/// handlers that declare the trailing argument: `visit` receives the state +/// active at the value and forwards it — or an override — when descending. +/// Reflected-field annotations override the forwarded state automatically +/// inside `default_visit_children`. +pub trait StructuralVisitor: Sized { + /// Visit one value under the definition-region state active at it. + fn visit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result>; + + /// Visit `child` now under `def_region_kind`, dispatching back into + /// [`StructuralVisitor::visit`]. An FFI `None` child is skipped without + /// a callback, matching the walk layer. + #[inline] + fn visit_child( + &mut self, + child: &T, + def_region_kind: DefRegionKind, + ) -> Result> + where + for<'x> AnyView<'x>: From<&'x T>, + { + let raw = raw_of(AnyView::from(child)); + if raw.type_index == TVMFFITypeIndex::kTVMFFINone as i32 { + return Ok(None); + } + self.visit(&VisitValue::from_raw(raw), def_region_kind) + } + + /// Visit `value`'s children — not `value` itself — with the default + /// rules, dispatching each child back into [`StructuralVisitor::visit`]. + /// + /// Children are container contents for `Array`/`List`/`Map`/`Dict` and + /// reflected structural fields otherwise. Field annotations override + /// `def_region_kind` for that field's recursive visit exactly like the + /// walk layer. + #[inline] + fn default_visit_children( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result> { + let result = visit_children_raw( + value.0, + &mut UserChildren { visitor: self }, + def_region_kind, + ) + .map_err(|halt| with_value_context(halt, value.0)); + finish(result) + } +} + +/// Internal per-value protocol driven by the recursion engine. Public only +/// as the bound of [`IntoWalker::Walker`]; not meant to be implemented +/// outside this crate. +#[doc(hidden)] +pub trait NativeVisit { + fn enter(&mut self, value: &VisitValue, def_region_kind: DefRegionKind) -> Result; + + fn exit(&mut self, _value: &VisitValue, _def_region_kind: DefRegionKind) -> Result { + Ok(WalkResult::Advance) + } +} + +/// Per-child action invoked by the shared child-iteration engine. +/// +/// The engine owns *finding* the children (container contents, reflected +/// fields) and computing each child's definition-region state; this trait +/// decides what happens at a child. The walk layer recurses +/// ([`WalkChildren`]); the visitor layer hands the child straight to user +/// code ([`UserChildren`]). +trait ChildVisit { + fn visit_child(&mut self, child: TVMFFIAny, def_region_kind: DefRegionKind) -> NativeResult; +} + +/// Walk-layer recursion: every child re-enters [`visit_raw`]. +struct WalkChildren<'a, V> { + visitor: &'a mut V, +} + +impl ChildVisit for WalkChildren<'_, V> { + fn visit_child(&mut self, child: TVMFFIAny, def_region_kind: DefRegionKind) -> NativeResult { + visit_raw(child, self.visitor, def_region_kind) + } +} + +/// Visitor-layer dispatch: every child goes back into the user-driven +/// [`StructuralVisitor::visit`], which controls further descent itself. +struct UserChildren<'a, V> { + visitor: &'a mut V, +} + +impl ChildVisit for UserChildren<'_, V> { + #[inline] + fn visit_child(&mut self, child: TVMFFIAny, def_region_kind: DefRegionKind) -> NativeResult { + if child.type_index == TVMFFITypeIndex::kTVMFFINone as i32 { + return Ok(()); + } + match self + .visitor + .visit(&VisitValue::from_raw(child), def_region_kind) + { + Ok(None) => Ok(()), + Ok(Some(interrupt)) => Err(NativeHalt::Interrupt(interrupt.value)), + Err(error) => Err(NativeHalt::Error(error)), + } + } +} + +/// Recurse into `value` on behalf of `visitor`: fire its enter hook, walk the +/// children, fire its exit hook. The engine below is stateless — these are +/// free functions, with the only shared piece (the `__s_visit__` attribute +/// column) cached process-wide. +fn visit_raw( + value: TVMFFIAny, + visitor: &mut V, + def_region_kind: DefRegionKind, +) -> NativeResult { + if value.type_index == TVMFFITypeIndex::kTVMFFINone as i32 { + return Ok(()); + } + + let visit_value = VisitValue::from_raw(value); + // Single by-value matches: splitting the Result match from the + // WalkResult match leaves a partially-moved temporary whose drop glue + // the compiler cannot fold away (measurably so on the container fast + // path). + match visitor.enter(&visit_value, def_region_kind) { + Ok(WalkResult::Advance) => {} + Ok(WalkResult::Skip) => return Ok(()), + Ok(WalkResult::Interrupt) => return Err(NativeHalt::Interrupt(Any::new())), + Ok(WalkResult::InterruptWith(payload)) => return Err(NativeHalt::Interrupt(payload)), + Err(error) => return Err(with_value_context(error.into(), value)), + } + + let children = &mut WalkChildren { + visitor: &mut *visitor, + }; + if let Err(halt) = visit_children_raw(value, children, def_region_kind) { + return Err(with_value_context(halt, value)); + } + + match visitor.exit(&visit_value, def_region_kind) { + Ok(WalkResult::Interrupt) => Err(NativeHalt::Interrupt(Any::new())), + Ok(WalkResult::InterruptWith(payload)) => Err(NativeHalt::Interrupt(payload)), + Ok(WalkResult::Advance | WalkResult::Skip) => Ok(()), + Err(error) => Err(with_value_context(error.into(), value)), + } +} + +#[inline] +fn visit_children_raw( + value: TVMFFIAny, + visitor: &mut C, + def_region_kind: DefRegionKind, +) -> NativeResult { + match value.type_index { + x if x == TVMFFITypeIndex::kTVMFFIArray as i32 + || x == TVMFFITypeIndex::kTVMFFIList as i32 => + { + return visit_sequence(value, visitor, def_region_kind); + } + x if x == TVMFFITypeIndex::kTVMFFIMap as i32 + || x == TVMFFITypeIndex::kTVMFFIDict as i32 => + { + // Fast path: read the MapBaseObj storage layout directly, like + // the SeqPrefix path for arrays — zero FFI calls per entry. + // Dict entries are snapshotted first to keep the re-entrant + // mutation guard. If the one-time layout validation fails + // (e.g. an ABI-debug build), fall back to the packed-functor + // iteration protocol. + if map_layout_usable(value) { + let snapshot = x == TVMFFITypeIndex::kTVMFFIDict as i32; + return visit_map_layout(value, visitor, def_region_kind, snapshot); + } + return visit_map(value, visitor, def_region_kind); + } + _ => {} + } + + reject_foreign_structural_visit(value.type_index)?; + if value.type_index < TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 { + Ok(()) + } else { + visit_reflected_fields(value, visitor, def_region_kind) + } +} + +#[inline(never)] +fn visit_sequence( + value: TVMFFIAny, + visitor: &mut C, + def_region_kind: DefRegionKind, +) -> NativeResult { + let seq = unsafe { &*(value.data_union.v_obj as *const SeqPrefix) }; + if seq.size < 0 { + return Err(runtime_error("native visitor: sequence reports a negative size").into()); + } + if seq.data.is_null() && seq.size != 0 { + return Err( + runtime_error("native visitor: non-empty sequence has a null data pointer").into(), + ); + } + let size = usize::try_from(seq.size) + .map_err(|_| runtime_error("native visitor: sequence size does not fit usize"))?; + if size == 0 { + return Ok(()); + } + + if value.type_index == TVMFFITypeIndex::kTVMFFIList as i32 { + // List storage may be invalidated by a re-entrant callback. Own a + // snapshot before running the first callback. + let children: Vec = { + let cells = unsafe { std::slice::from_raw_parts(seq.data, size) }; + cells + .iter() + .map(|cell| Any::from(unsafe { view_of(cell) })) + .collect() + }; + for (index, child) in children.into_iter().enumerate() { + let raw = raw_of_owned(&child); + visitor + .visit_child(raw, def_region_kind) + .map_err(|halt| with_error_context(halt, &format!("sequence item [{index}]")))?; + } + return Ok(()); + } + + // Array is immutable, so its element cells remain stable throughout + // recursive callbacks and need no refcounted snapshot. + let cells = unsafe { std::slice::from_raw_parts(seq.data, size) }; + for (index, child) in cells.iter().enumerate() { + visitor + .visit_child(*child, def_region_kind) + .map_err(|halt| with_error_context(halt, &format!("sequence item [{index}]")))?; + } + Ok(()) +} + +/// Walk map/dict entries by reading the `MapBaseObj` storage directly — +/// the map analog of the `SeqPrefix` array fast path. `snapshot` first +/// takes owned copies of all entries (Dict re-entrant mutation guard). +#[inline(never)] +fn visit_map_layout( + value: TVMFFIAny, + visitor: &mut C, + def_region_kind: DefRegionKind, + snapshot: bool, +) -> NativeResult { + let map = unsafe { &*(value.data_union.v_obj as *const MapPrefix) }; + let size = map.size as usize; + if size == 0 { + return Ok(()); + } + let mut cursor = unsafe { MapCursor::new(map) }; + + if snapshot { + let mut entries: Vec<(Any, Any)> = Vec::with_capacity(size); + for _ in 0..size { + let Some((key, val)) = (unsafe { cursor.next() }) else { + return Err(runtime_error("native visitor: map iteration ended early").into()); + }; + entries.push(( + Any::from(unsafe { view_of(&key) }), + Any::from(unsafe { view_of(&val) }), + )); + } + for (index, (key, val)) in entries.into_iter().enumerate() { + visitor + .visit_child(raw_of_owned(&key), def_region_kind) + .map_err(|halt| with_error_context(halt, &format!("dict key [{index}]")))?; + visitor + .visit_child(raw_of_owned(&val), def_region_kind) + .map_err(|halt| with_error_context(halt, &format!("dict value [{index}]")))?; + } + return Ok(()); + } + + // Immutable map: entry cells stay stable throughout recursive + // callbacks, so visit them in place. The `size` bound also guards the + // dense iteration list against corruption-induced cycles. + for index in 0..size { + let Some((key, val)) = (unsafe { cursor.next() }) else { + return Err(runtime_error("native visitor: map iteration ended early").into()); + }; + visitor + .visit_child(key, def_region_kind) + .map_err(|halt| with_error_context(halt, &format!("map key [{index}]")))?; + visitor + .visit_child(val, def_region_kind) + .map_err(|halt| with_error_context(halt, &format!("map value [{index}]")))?; + } + Ok(()) +} + +/// Cold fallback used when the mirrored layout fails validation (e.g. an +/// ABI-debug build): iterate through the public packed functors. Map storage +/// is private C++; the Rust binding itself uses these iterator functors, so +/// no structural visiting or traversal control leaves Rust. Entries are +/// snapshotted before user callbacks run — required for Dict, whose mutation +/// invalidates the iterator, and harmless for immutable Map on this +/// non-performance path. +fn visit_map( + value: TVMFFIAny, + visitor: &mut C, + def_region_kind: DefRegionKind, +) -> NativeResult { + let is_dict = value.type_index == TVMFFITypeIndex::kTVMFFIDict as i32; + let (size_name, iter_name, kind) = if is_dict { + ("ffi.DictSize", "ffi.DictForwardIterFunctor", "dict") + } else { + ("ffi.MapSize", "ffi.MapForwardIterFunctor", "map") + }; + let size = Function::get_global(size_name)? + .call_packed(&[unsafe { view_of(&value) }]) + .and_then(i64::try_from)?; + if size < 0 { + return Err(runtime_error("native visitor: map reports a negative size").into()); + } + let size = usize::try_from(size) + .map_err(|_| runtime_error("native visitor: map size does not fit usize"))?; + if size == 0 { + return Ok(()); + } + + let iter_any = Function::get_global(iter_name)?.call_packed(&[unsafe { view_of(&value) }])?; + let iter = Function::try_from(iter_any)?; + + let mut entries = Vec::with_capacity(size); + for index in 0..size { + let key = iter.call_packed(&[AnyView::from(&0i64)])?; + let map_value = iter.call_packed(&[AnyView::from(&1i64)])?; + entries.push((key, map_value)); + if index + 1 != size { + iter.call_packed(&[AnyView::from(&2i64)])?; + } + } + + for (index, (key, map_value)) in entries.into_iter().enumerate() { + visitor + .visit_child(raw_of_owned(&key), def_region_kind) + .map_err(|halt| with_error_context(halt, &format!("{kind} key [{index}]")))?; + visitor + .visit_child(raw_of_owned(&map_value), def_region_kind) + .map_err(|halt| with_error_context(halt, &format!("{kind} value [{index}]")))?; + } + Ok(()) +} + +#[inline] +fn visit_reflected_fields( + value: TVMFFIAny, + visitor: &mut C, + def_region_kind: DefRegionKind, +) -> NativeResult { + let type_info = unsafe { TVMFFIGetTypeInfo(value.type_index) }; + if type_info.is_null() { + return Err(runtime_error(&format!( + "native visitor: unregistered type index {}", + value.type_index + )) + .into()); + } + let seq_hash_kind = unsafe { + let metadata = (*type_info).metadata; + if metadata.is_null() { + TVMFFISEqHashKind::kTVMFFISEqHashKindUnsupported as i32 + } else { + (*metadata).structural_eq_hash_kind + } + }; + let def_region_kind = free_var_child_region(def_region_kind, seq_hash_kind); + let object = unsafe { value.data_union.v_obj } as *mut u8; + let halted = unsafe { + for_each_field(value.type_index, |field| { + match visit_reflected_field(object, field, visitor, def_region_kind) { + Ok(()) => ControlFlow::Continue(()), + Err(halt) => ControlFlow::Break(halt), + } + }) + }; + halted.map_or(Ok(()), Err) +} + +unsafe fn visit_reflected_field( + object: *mut u8, + field: &TVMFFIFieldInfo, + visitor: &mut C, + inherited_region: DefRegionKind, +) -> NativeResult { + if field.flags & FLAG_SEQ_HASH_IGNORE != 0 { + return Ok(()); + } + + let Some(getter) = field.getter else { + return Err(NativeHalt::Error(runtime_error(&format!( + "native visitor: reflected field `{}` has no getter", + field.name.as_str() + )))); + }; + let address = object.offset(field.offset as isize) as *mut c_void; + let mut child_raw = TVMFFIAny::new(); + if getter(address, &mut child_raw) != 0 { + return Err(with_error_context( + NativeHalt::Error(Error::from_raised()), + &format!("field `{}`", field.name.as_str()), + )); + } + + // A reflection getter returns an owned Any. Keep it alive while the + // recursive walk borrows its raw cell. + let child = Any::from_raw_ffi_any(child_raw); + let borrowed = raw_of_owned(&child); + let child_region = field_def_region(field, inherited_region); + visitor + .visit_child(borrowed, child_region) + .map_err(|halt| with_error_context(halt, &format!("field `{}`", field.name.as_str()))) +} + +// Runs once per visited value: keep the no-hook fast path small enough to +// actually inline (one cached-column load and a tag compare) and the error +// formatting out of line — with the cold body inside, the `#[inline]` hint +// was declined and the call cost ~20% of the container fast path. +#[inline] +fn reject_foreign_structural_visit(type_index: i32) -> Result<()> { + let Some(attr) = structural_visit_column().and_then(|column| column.get(type_index)) else { + return Ok(()); + }; + if attr.type_index == TVMFFITypeIndex::kTVMFFINone as i32 { + return Ok(()); + } + reject_foreign_structural_visit_cold(type_index, attr.type_index) +} + +#[cold] +#[inline(never)] +fn reject_foreign_structural_visit_cold(type_index: i32, attr_type_index: i32) -> Result<()> { + if attr_type_index == TVMFFITypeIndex::kTVMFFIOpaquePtr as i32 + || attr_type_index == TVMFFITypeIndex::kTVMFFIFunction as i32 + { + let value_type = if type_index < TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 { + format!("type index {type_index}") + } else { + format!("type `{}`", type_key_of(type_index)) + }; + Err(runtime_error(&format!( + "native visitor: {value_type} registers foreign `{STRUCTURAL_VISIT_ATTR}`; \ + visit its children explicitly from a `StructuralVisitor` \ + (`structural_visit`), or skip it with a pre-order `WalkResult::Skip` \ + handler" + ))) + } else { + Err(Error::new( + TYPE_ERROR, + &format!("{STRUCTURAL_VISIT_ATTR} must be an opaque function pointer or ffi.Function"), + "", + )) + } +} + +fn with_value_context(halt: NativeHalt, value: TVMFFIAny) -> NativeHalt { + if value.type_index < TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 { + halt + } else { + with_error_context(halt, &format!("object `{}`", type_key_of(value.type_index))) + } +} + +/// Visit `root` with a user-driven [`StructuralVisitor`]. +/// +/// The visitor's [`StructuralVisitor::visit`] runs for the root under +/// [`DefRegionKind::None`] and controls all further recursion itself. This is +/// the Rust analog of constructing a C++ `StructuralVisitorObj` and calling +/// `visitor->Visit(root)`. An FFI `None` root completes immediately. +pub fn structural_visit(root: &R, visitor: &mut V) -> Result> +where + V: StructuralVisitor, + for<'x> AnyView<'x>: From<&'x R>, +{ + visitor.visit_child(root, DefRegionKind::None) +} + +/// Walk `root` with an observer, the Rust analog of C++ +/// `StructuralWalk(root, callbacks...)`. +/// +/// `walker` is anything implementing [`IntoWalker`]: a `&mut` reference to a +/// stateful [`VisitDispatch`] visitor (`#[dispatch(visit)]`), a bare closure +/// in any [`WalkChainLink`] shape (catch-all `&VisitValue`, typed, or node, +/// with an optional trailing [`DefRegionKind`]), or a tuple of such +/// callbacks tried in order — the C++ callback overloads and variadic +/// chain. The walker owns recursion: the handler runs once per value, +/// before or after the value's children according to `order`, and steers +/// traversal through the returned [`WalkResult`]. +pub fn structural_walk( + root: &R, + walker: H, + order: WalkOrder, +) -> Result> +where + H: IntoWalker, + for<'x> AnyView<'x>: From<&'x R>, +{ + let mut dispatch = walker.into_walker(order); + finish(visit_raw( + raw_of(AnyView::from(root)), + &mut dispatch, + DefRegionKind::None, + )) +} + +fn finish(result: NativeResult) -> Result> { + match result { + Ok(()) => Ok(None), + Err(NativeHalt::Error(error)) => Err(error), + Err(NativeHalt::Interrupt(payload)) => Ok(Some(VisitInterrupt { value: payload })), + } +} + +#[inline] +fn field_def_region(field: &TVMFFIFieldInfo, inherited: DefRegionKind) -> DefRegionKind { + if field.flags & FLAG_SEQ_HASH_DEF_NON_RECURSIVE != 0 { + DefRegionKind::NonRecursive + } else if field.flags & FLAG_SEQ_HASH_DEF_RECURSIVE != 0 { + DefRegionKind::Recursive + } else { + inherited + } +} + +/// A non-recursive definition applies to a FreeVar value itself, but not to +/// the FreeVar's own reflected children: nested free vars there must resolve +/// against an outer binding instead of rebinding. Mirrors C++ +/// `VisitReflectedFieldsExpected`. +#[inline] +fn free_var_child_region(inherited: DefRegionKind, structural_eq_hash_kind: i32) -> DefRegionKind { + if inherited == DefRegionKind::NonRecursive + && structural_eq_hash_kind == TVMFFISEqHashKind::kTVMFFISEqHashKindFreeVar as i32 + { + DefRegionKind::None + } else { + inherited + } +} + +fn with_error_context(halt: NativeHalt, frame: &str) -> NativeHalt { + match halt { + NativeHalt::Error(error) => NativeHalt::Error(Error::with_appended_backtrace( + error, + &format!("[native structural visit] {frame}\n"), + )), + interrupt => interrupt, + } +} + +fn runtime_error(message: &str) -> Error { + Error::new(RUNTIME_ERROR, message, "") +} + +/// Layout prefix shared by the C++ `MapObj` and `DictObj` (`MapBaseObj`, +/// release ABI without `TVM_FFI_DEBUG_WITH_ABI_CHANGE`). +#[repr(C)] +struct MapPrefix { + _header: TVMFFIObject, + data: *mut u8, + size: u64, + slots: u64, + _data_deleter: Option, +} + +/// Dense-layout extension of the prefix (`DenseMapBaseObj`). +#[repr(C)] +struct DenseMapPrefix { + base: MapPrefix, + fib_shift: u32, + iter_list_head: u64, + iter_list_tail: u64, +} + +const _: () = { + assert!(std::mem::offset_of!(MapPrefix, data) == 24); + assert!(std::mem::offset_of!(MapPrefix, size) == 32); + assert!(std::mem::offset_of!(MapPrefix, slots) == 40); + assert!(std::mem::offset_of!(MapPrefix, _data_deleter) == 48); + assert!(std::mem::offset_of!(DenseMapPrefix, fib_shift) == 56); + assert!(std::mem::offset_of!(DenseMapPrefix, iter_list_head) == 64); +}; + +/// MSB tag on `slots_` marking the small (inline KV array) layout. +const MAP_SMALL_TAG: u64 = 1 << 63; +/// `kInvalidIndex`: terminator of the dense iteration list. +const MAP_INVALID_INDEX: u64 = u64::MAX; +/// `kBlockCap`: entries per dense block. +const MAP_BLOCK_CAP: u64 = 16; +/// `sizeof(ItemType)`: KV pair (32 bytes) + prev/next indices (16 bytes). +const MAP_ITEM_SIZE: usize = 48; +/// `sizeof(Block)`: `kBlockCap` metadata bytes + `kBlockCap` items. +const MAP_BLOCK_SIZE: usize = 16 + 16 * MAP_ITEM_SIZE; +/// Byte offset of `ItemType::next` (after the 32-byte KV pair and `prev`). +const MAP_ITEM_NEXT_OFFSET: usize = 40; + +/// Borrowed traversal cursor over either map storage layout, yielding entries +/// in the same order as the C++ iterator. +enum MapCursor { + Small { + kv: *const TVMFFIAny, + index: usize, + size: usize, + }, + Dense { + data: *const u8, + index: u64, + }, +} + +impl MapCursor { + #[inline] + unsafe fn new(map: &MapPrefix) -> MapCursor { + if map.slots & MAP_SMALL_TAG != 0 { + MapCursor::Small { + kv: map.data as *const TVMFFIAny, + index: 0, + size: map.size as usize, + } + } else { + let dense = &*(map as *const MapPrefix as *const DenseMapPrefix); + MapCursor::Dense { + data: map.data, + index: dense.iter_list_head, + } + } + } + + #[inline] + unsafe fn next(&mut self) -> Option<(TVMFFIAny, TVMFFIAny)> { + match self { + MapCursor::Small { kv, index, size } => { + if *index >= *size { + return None; + } + let pair = kv.add(*index * 2); + *index += 1; + Some((*pair, *pair.add(1))) + } + MapCursor::Dense { data, index } => { + if *index == MAP_INVALID_INDEX { + return None; + } + let block = data.add((*index / MAP_BLOCK_CAP) as usize * MAP_BLOCK_SIZE); + let item = block.add( + MAP_BLOCK_CAP as usize + (*index % MAP_BLOCK_CAP) as usize * MAP_ITEM_SIZE, + ); + let key = *(item as *const TVMFFIAny); + let val = *(item.add(16) as *const TVMFFIAny); + *index = *(item.add(MAP_ITEM_NEXT_OFFSET) as *const u64); + Some((key, val)) + } + } + } +} + +/// Process-wide result of the one-time map layout validation: +/// 0 = unknown, 1 = usable, 2 = unusable. +static MAP_LAYOUT_STATE: AtomicU8 = AtomicU8::new(0); + +#[inline] +fn map_layout_usable(value: TVMFFIAny) -> bool { + match MAP_LAYOUT_STATE.load(Ordering::Relaxed) { + 1 => true, + 2 => false, + _ => { + let usable = validate_map_layout(value); + MAP_LAYOUT_STATE.store(if usable { 1 } else { 2 }, Ordering::Relaxed); + usable + } + } +} + +/// Cross-check the mirrored `MapBaseObj` layout against the public size +/// functor once per process. An ABI-debug build inserts a state marker that +/// shifts every field by 8 bytes, which this detects: offset 32 then holds a +/// pointer value that cannot equal the reported entry count. +fn validate_map_layout(value: TVMFFIAny) -> bool { + let expected = (|| -> Result { + let is_dict = value.type_index == TVMFFITypeIndex::kTVMFFIDict as i32; + let name = if is_dict { + "ffi.DictSize" + } else { + "ffi.MapSize" + }; + Function::get_global(name)? + .call_packed(&[unsafe { view_of(&value) }]) + .and_then(i64::try_from) + })(); + let Ok(expected) = expected else { + return false; + }; + let map = unsafe { &*(value.data_union.v_obj as *const MapPrefix) }; + expected >= 0 && map.size == expected as u64 +} + +/// Layout prefix shared by the C++ `ArrayObj` and `ListObj`. +#[repr(C)] +struct SeqPrefix { + _header: TVMFFIObject, + data: *const TVMFFIAny, + size: i64, +} + +const _: () = { + assert!(std::mem::offset_of!(SeqPrefix, data) == 24); + assert!(std::mem::offset_of!(SeqPrefix, size) == 32); +}; + +#[derive(Clone, Copy)] +struct TypeAttrColumn(NonNull); + +impl TypeAttrColumn { + /// Copy one borrowed cell; ownership remains with the registry. + fn get(self, type_index: i32) -> Option { + unsafe { + let column = self.0.as_ref(); + let index = type_index - column.begin_index; + if index < 0 || index >= column.size || column.data.is_null() { + None + } else { + Some(*column.data.offset(index as isize)) + } + } + } +} + +fn type_attr_column(attr_name: &str) -> Option { + unsafe { + let attr_name = TVMFFIByteArray::from_str(attr_name); + NonNull::new(TVMFFIGetTypeAttrColumn(&attr_name).cast_mut()).map(TypeAttrColumn) + } +} + +/// Cached `__s_visit__` column pointer (0 = not seen yet). A registry column +/// is stable once created — C++ `DefaultVisitExpected` caches the same +/// pointer in a function-local static — while an absent column is re-queried +/// because a later attr registration may create it. The cache keeps the +/// per-value foreign-hook check free of FFI lookups. +static STRUCTURAL_VISIT_COLUMN: AtomicUsize = AtomicUsize::new(0); + +#[inline] +fn structural_visit_column() -> Option { + let cached = STRUCTURAL_VISIT_COLUMN.load(Ordering::Relaxed); + if cached != 0 { + let pointer = cached as *mut TVMFFITypeAttrColumn; + return Some(TypeAttrColumn(unsafe { NonNull::new_unchecked(pointer) })); + } + let column = type_attr_column(STRUCTURAL_VISIT_ATTR)?; + STRUCTURAL_VISIT_COLUMN.store(column.0.as_ptr() as usize, Ordering::Relaxed); + Some(column) +} + +fn type_key_of(type_index: i32) -> String { + unsafe { + let info = TVMFFIGetTypeInfo(type_index); + if info.is_null() { + format!("") + } else { + (*info).type_key.as_str().to_string() + } + } +} + +/// Subtype check with the base's inheritance depth supplied by the caller +/// (`ObjectCore::TYPE_DEPTH`), so only the object's type info is fetched. +#[inline] +fn is_instance_at_depth(object_type_index: i32, base_type_index: i32, base_depth: i32) -> bool { + if object_type_index == base_type_index { + return true; + } + unsafe { + let info = TVMFFIGetTypeInfo(object_type_index); + if info.is_null() { + return false; + } + if (*info).type_depth <= base_depth { + return false; + } + let ancestors = (*info).type_acenstors; + if ancestors.is_null() { + return false; + } + let ancestor = *ancestors.offset(base_depth as isize); + !ancestor.is_null() && (*ancestor).type_index == base_type_index + } +} + +/// Visit every reflected field of `type_index` and its ancestors in the same +/// parent-to-child order as C++ `ForEachFieldInfoWithEarlyStop`. +/// +/// # Safety +/// +/// `type_index` must be a registered type index. +unsafe fn for_each_field( + type_index: i32, + mut callback: impl FnMut(&'static TVMFFIFieldInfo) -> ControlFlow, +) -> Option { + let info = TVMFFIGetTypeInfo(type_index); + if info.is_null() { + return None; + } + + // Ancestor slot 0 is the root Object. C++ starts at slot 1, walks toward + // the immediate parent, then visits the concrete type's own fields. + for depth in 1..(*info).type_depth { + let ancestor = *(*info).type_acenstors.offset(depth as isize); + if let Some(value) = visit_field_level(ancestor, &mut callback) { + return Some(value); + } + } + visit_field_level(info, &mut callback) +} + +unsafe fn visit_field_level( + info: *const crate::tvm_ffi_sys::TVMFFITypeInfo, + callback: &mut impl FnMut(&'static TVMFFIFieldInfo) -> ControlFlow, +) -> Option { + if info.is_null() || (*info).fields.is_null() { + return None; + } + let fields = std::slice::from_raw_parts((*info).fields, (*info).num_fields as usize); + for field in fields { + // C reflection tables are immortal once registered. + let field: &'static TVMFFIFieldInfo = &*(field as *const TVMFFIFieldInfo); + if let ControlFlow::Break(value) = callback(field) { + return Some(value); + } + } + None +} + +#[inline] +fn raw_of(view: AnyView<'_>) -> TVMFFIAny { + *view.as_raw_ffi_any() +} + +#[inline] +fn raw_of_owned(any: &Any) -> TVMFFIAny { + *any.as_raw_ffi_any() +} + +#[inline] +unsafe fn view_of(raw: &TVMFFIAny) -> AnyView<'_> { + unsafe { AnyView::from_raw_ffi_any(*raw) } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::Array; + + struct RegionProbe(Vec); + + impl NativeVisit for RegionProbe { + fn enter( + &mut self, + _value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result { + self.0.push(def_region_kind); + Ok(WalkResult::Advance) + } + } + + #[derive(Default)] + struct TypedRegionProbe { + seen: Vec, + } + + #[crate::dispatch(visit)] + impl TypedRegionProbe { + fn visit_integer(&mut self, _value: i64, def_region_kind: DefRegionKind) -> WalkResult { + self.seen.push(def_region_kind); + WalkResult::Advance + } + } + + unsafe extern "C" fn clone_any_field(field: *mut c_void, result: *mut TVMFFIAny) -> i32 { + let value = &*(field as *const Any); + *result = Any::into_raw_ffi_any(value.clone()); + 0 + } + + #[test] + fn def_region_is_inherited_through_containers() { + let root = Array::new(vec![1i64, 2]); + let mut probe = RegionProbe(Vec::new()); + assert!(visit_raw( + raw_of(AnyView::from(&root)), + &mut probe, + DefRegionKind::Recursive, + ) + .is_ok()); + assert_eq!(probe.0, vec![DefRegionKind::Recursive; 3]); + } + + #[test] + fn reflected_field_def_region_reaches_typed_handler() { + let mut probe = TypedRegionProbe::default(); + let mut dispatch = (&mut probe).into_walker(WalkOrder::PreOrder); + let mut value = Any::from(7i64); + let mut field: TVMFFIFieldInfo = unsafe { std::mem::zeroed() }; + field.name = unsafe { TVMFFIByteArray::from_str("value") }; + field.getter = Some(clone_any_field); + let object = (&mut value as *mut Any).cast::(); + + let mut children = WalkChildren { + visitor: &mut dispatch, + }; + for flags in [ + FLAG_SEQ_HASH_DEF_RECURSIVE, + 0, + FLAG_SEQ_HASH_DEF_NON_RECURSIVE, + FLAG_SEQ_HASH_DEF_NON_RECURSIVE | FLAG_SEQ_HASH_DEF_RECURSIVE, + FLAG_SEQ_HASH_IGNORE, + ] { + field.flags = flags; + assert!(unsafe { + visit_reflected_field(object, &field, &mut children, DefRegionKind::None) + } + .is_ok()); + } + assert_eq!( + probe.seen, + vec![ + DefRegionKind::Recursive, + DefRegionKind::None, + DefRegionKind::NonRecursive, + DefRegionKind::NonRecursive, + ] + ); + } + + #[test] + fn non_recursive_region_is_clamped_for_free_var_children_only() { + use TVMFFISEqHashKind::{kTVMFFISEqHashKindFreeVar, kTVMFFISEqHashKindTreeNode}; + + let free_var = kTVMFFISEqHashKindFreeVar as i32; + let tree_node = kTVMFFISEqHashKindTreeNode as i32; + assert_eq!( + free_var_child_region(DefRegionKind::NonRecursive, free_var), + DefRegionKind::None + ); + assert_eq!( + free_var_child_region(DefRegionKind::Recursive, free_var), + DefRegionKind::Recursive + ); + assert_eq!( + free_var_child_region(DefRegionKind::None, free_var), + DefRegionKind::None + ); + assert_eq!( + free_var_child_region(DefRegionKind::NonRecursive, tree_node), + DefRegionKind::NonRecursive + ); + } +} diff --git a/rust/tvm-ffi/src/lib.rs b/rust/tvm-ffi/src/lib.rs index 124923080..865798576 100644 --- a/rust/tvm-ffi/src/lib.rs +++ b/rust/tvm-ffi/src/lib.rs @@ -46,13 +46,18 @@ pub use crate::error::{ ATTRIBUTE_ERROR, INDEX_ERROR, KEY_ERROR, RUNTIME_ERROR, TYPE_ERROR, VALUE_ERROR, }; pub use crate::extra::module::Module; +pub use crate::extra::structural_visit::{ + structural_visit, structural_walk, DefRegionKind, IntoVisitResult, IntoWalker, + StructuralVisitor, VisitDispatch, VisitInterrupt, VisitValue, WalkChainLink, WalkOrder, + WalkResult, +}; pub use crate::function::Function; pub use crate::object::ObjectRefCast; pub use crate::object::{Object, ObjectArc, ObjectCore, ObjectCoreWithExtraItems, ObjectRefCore}; pub use crate::optional::Optional; pub use crate::string::{Bytes, String}; pub use crate::type_traits::AnyCompatible; -pub use tvm_ffi_macros::match_any; +pub use tvm_ffi_macros::{dispatch, match_any}; pub use tvm_ffi_sys::TVMFFITypeIndex as TypeIndex; pub use tvm_ffi_sys::{ diff --git a/rust/tvm-ffi/tests/test_structural_visit.rs b/rust/tvm-ffi/tests/test_structural_visit.rs new file mode 100644 index 000000000..723de8079 --- /dev/null +++ b/rust/tvm-ffi/tests/test_structural_visit.rs @@ -0,0 +1,873 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +use tvm_ffi::tvm_ffi_sys::{TVMFFIByteArray, TVMFFITypeIndex, TVMFFITypeRegisterAttr}; +use tvm_ffi::{ + dispatch, structural_visit, structural_walk, Any, AnyView, Array, DefRegionKind, Error, + Function, Map, Object, Result, Shape, String as FfiString, StructuralVisitor, VisitInterrupt, + VisitValue, WalkOrder, WalkResult, RUNTIME_ERROR, +}; + +fn runtime_error(message: &str) -> Error { + Error::new(RUNTIME_ERROR, message, "") +} + +#[test] +fn plain_walk_uses_native_sequence_fallback() { + let root = Array::new(vec![1i64, 2, 3]); + let mut integers = 0; + assert!(structural_walk( + &root, + |value: &VisitValue| { + if value.cast::().is_some() { + integers += 1; + } + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!(integers, 3); +} + +#[test] +fn plain_walk_uses_native_map_fallback() { + let root: Map = [(FfiString::from("a"), 1i64), (FfiString::from("b"), 2i64)] + .into_iter() + .collect(); + let mut integers = 0; + assert!(structural_walk( + &root, + |value: &VisitValue| { + if value.cast::().is_some() { + integers += 1; + } + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!(integers, 2); +} + +#[derive(Default)] +struct SkipForeignShape {} + +#[dispatch(visit)] +impl SkipForeignShape { + fn visit_shape(&mut self, _shape: Shape) -> WalkResult { + WalkResult::Skip + } +} + +/// Visitor-layer handling of the foreign type: `visit` enumerates the +/// children itself (none, for a shape) instead of the default recursion. +#[derive(Default)] +struct ForeignShapeVisitor { + shapes: usize, +} + +impl StructuralVisitor for ForeignShapeVisitor { + fn visit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result> { + if value.cast::().is_some() { + self.shapes += 1; + return Ok(None); + } + self.default_visit_children(value, def_region_kind) + } +} + +#[test] +fn foreign_structural_visit_requires_explicit_rust_override() { + let hook = Function::get_global("ffi.ArraySize").unwrap(); + let attr_name = unsafe { TVMFFIByteArray::from_str("__s_visit__") }; + let mut attr_value = Any::from(hook); + assert_eq!( + unsafe { + TVMFFITypeRegisterAttr( + TVMFFITypeIndex::kTVMFFIShape as i32, + &attr_name, + Any::as_data_ptr(&mut attr_value), + ) + }, + 0 + ); + + let root = Shape::from([2i64, 3]); + let error = match structural_walk( + &root, + |_value: &VisitValue| WalkResult::Advance, + WalkOrder::PreOrder, + ) { + Err(error) => error, + Ok(_) => panic!("foreign structural visit unexpectedly used reflection"), + }; + assert!(error.message().contains("registers foreign `__s_visit__`")); + assert!(error.message().contains("StructuralVisitor")); + + // Walk layer: a pre-order handler skips the foreign type. + assert!( + structural_walk(&root, &mut SkipForeignShape::default(), WalkOrder::PreOrder) + .unwrap() + .is_none() + ); + + // Visitor layer: take over the type's children explicitly instead. + let mut takeover = ForeignShapeVisitor::default(); + assert!(structural_visit(&root, &mut takeover).unwrap().is_none()); + assert_eq!(takeover.shapes, 1); +} + +#[test] +fn mutable_list_is_snapshotted_before_callbacks() { + let root = Function::get_global("ffi.List") + .unwrap() + .call_packed(&[AnyView::from(&1i64), AnyView::from(&2i64)]) + .unwrap(); + let captured = root.clone(); + let append = Function::get_global("ffi.ListAppend").unwrap(); + let mut appended = false; + let mut integers = Vec::new(); + + assert!(structural_walk( + &root, + |value: &VisitValue| { + if let Some(integer) = value.cast::() { + integers.push(integer); + if !appended { + append + .call_packed(&[AnyView::from(&captured), AnyView::from(&3i64)]) + .unwrap(); + appended = true; + } + } + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + + assert_eq!(integers, vec![1, 2]); + let size = Function::get_global("ffi.ListSize") + .unwrap() + .call_packed(&[AnyView::from(&root)]) + .and_then(i64::try_from) + .unwrap(); + assert_eq!(size, 3); +} + +#[test] +fn mutable_dict_is_snapshotted_before_callbacks() { + let root = Function::get_global("ffi.Dict") + .unwrap() + .call_packed(&[ + AnyView::from(&FfiString::from("a")), + AnyView::from(&1i64), + AnyView::from(&FfiString::from("b")), + AnyView::from(&2i64), + ]) + .unwrap(); + let captured = root.clone(); + let set_item = Function::get_global("ffi.DictSetItem").unwrap(); + let mut inserted = false; + let mut integers = Vec::new(); + + assert!(structural_walk( + &root, + |value: &VisitValue| { + if let Some(integer) = value.cast::() { + integers.push(integer); + if !inserted { + set_item + .call_packed(&[ + AnyView::from(&captured), + AnyView::from(&FfiString::from("c")), + AnyView::from(&3i64), + ]) + .unwrap(); + inserted = true; + } + } + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + + integers.sort_unstable(); + assert_eq!(integers, vec![1, 2]); + let size = Function::get_global("ffi.DictSize") + .unwrap() + .call_packed(&[AnyView::from(&root)]) + .and_then(i64::try_from) + .unwrap(); + assert_eq!(size, 3); +} + +#[test] +fn dense_map_layout_is_traversed_completely() { + // More than 4 entries forces the dense (block + iteration list) layout. + let root: Map = (0..9) + .map(|i| (FfiString::from(format!("k{i}")), i as i64)) + .collect(); + let mut sum = 0; + let mut strings = 0; + assert!(structural_walk( + &root, + |value: &VisitValue| { + if let Some(integer) = value.cast::() { + sum += integer; + } else if value.cast::().is_some() { + strings += 1; + } + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!(sum, (0..9).sum::()); + assert_eq!(strings, 9); +} + +#[test] +fn interrupt_payload_crosses_map_traversal() { + let root: Map = [(FfiString::from("a"), 1i64), (FfiString::from("b"), 2i64)] + .into_iter() + .collect(); + let outcome = structural_walk( + &root, + |value: &VisitValue| { + if value.cast::().is_some() { + return WalkResult::interrupt_with(99i64); + } + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap(); + let Some(interrupt) = outcome else { + panic!("map walk unexpectedly completed"); + }; + assert_eq!(i64::try_from(interrupt.value).unwrap(), 99); +} + +#[test] +fn handler_error_crosses_map_traversal() { + let root: Map = [(FfiString::from("a"), 1i64)].into_iter().collect(); + let error = match structural_walk( + &root, + |value: &VisitValue| -> Result { + if value.cast::().is_some() { + Err(runtime_error("map handler failed")) + } else { + Ok(WalkResult::Advance) + } + }, + WalkOrder::PreOrder, + ) { + Err(error) => error, + Ok(_) => panic!("map handler unexpectedly succeeded"), + }; + assert_eq!(error.message(), "map handler failed"); + assert!(error.backtrace().contains("object `ffi.Map`")); +} + +#[test] +fn interrupt_stops_without_running_remaining_callbacks() { + let root = Array::new(vec![1i64, 2, 3]); + let mut integers = 0; + let outcome = structural_walk( + &root, + |value: &VisitValue| { + if value.cast::().is_some() { + integers += 1; + return WalkResult::Interrupt; + } + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap(); + assert!(outcome.is_some()); + assert_eq!(integers, 1); +} + +/// Visitor-layer traversal that overrides the def-region for one child and +/// inherits it for the next, mirroring a C++ visitor using +/// `WithDefRegionKind`. +#[derive(Default)] +struct ManualRegionVisitor { + seen: Vec, +} + +impl StructuralVisitor for ManualRegionVisitor { + fn visit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result> { + if let Some(array) = value.cast::>() { + // Override the state for exactly this child's subtree... + let overridden = array.get(0).unwrap(); + if let Some(interrupt) = self.visit_child(&overridden, DefRegionKind::NonRecursive)? { + return Ok(Some(interrupt)); + } + // ...and forward the received state to inherit it. + let inherited = array.get(1).unwrap(); + return self.visit_child(&inherited, def_region_kind); + } + if value.cast::().is_some() { + self.seen.push(def_region_kind); + } + Ok(None) + } +} + +#[test] +fn manual_child_visit_can_override_def_region() { + let root = Array::new(vec![7i64, 8]); + let mut probe = ManualRegionVisitor::default(); + assert!(structural_visit(&root, &mut probe).unwrap().is_none()); + assert_eq!( + probe.seen, + vec![DefRegionKind::NonRecursive, DefRegionKind::None] + ); +} + +#[derive(Default)] +struct GenericDispatchProbe { + integers: Vec, + objects: usize, + catch_all: usize, +} + +#[dispatch(visit)] +impl GenericDispatchProbe { + fn visit_integer(&mut self, value: i64) -> WalkResult { + self.integers.push(value); + WalkResult::Advance + } + + // Trailing DefRegionKind: handlers may mix arities within one impl. + fn visit_object(&mut self, _value: &tvm_ffi::Object, kind: DefRegionKind) -> WalkResult { + assert_eq!(kind, DefRegionKind::None); + self.objects += 1; + WalkResult::Advance + } + + fn visit_any(&mut self, _value: &VisitValue) -> WalkResult { + self.catch_all += 1; + WalkResult::Advance + } +} + +#[test] +fn generated_dispatch_supports_pod_and_ordered_catch_all() { + let root = Array::new(vec![1i64, 2]); + let mut probe = GenericDispatchProbe::default(); + assert!(structural_walk(&root, &mut probe, WalkOrder::PreOrder) + .unwrap() + .is_none()); + assert_eq!(probe.integers, vec![1, 2]); + assert_eq!(probe.objects, 1); + + let floats = Array::new(vec![1.0f64, 2.0]); + assert!(structural_walk(&floats, &mut probe, WalkOrder::PreOrder) + .unwrap() + .is_none()); + assert_eq!(probe.objects, 2); + assert_eq!(probe.catch_all, 2); +} + +/// Visitor-layer enter/exit straddling: run enter logic, delegate the +/// default child recursion, then run exit logic with the same locals in +/// scope — the C++ `DefaultVisitExpected` pattern. +#[derive(Default)] +struct StraddleVisitor { + events: Vec, +} + +impl StructuralVisitor for StraddleVisitor { + fn visit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result> { + let label = match value.cast::() { + Some(integer) => format!("int:{integer}"), + None => "node".to_string(), + }; + self.events.push(format!("enter:{label}")); + if let Some(interrupt) = self.default_visit_children(value, def_region_kind)? { + return Ok(Some(interrupt)); + } + self.events.push(format!("exit:{label}")); + Ok(None) + } +} + +#[test] +fn visitor_can_straddle_default_children() { + let root = Array::new(vec![1i64, 2]); + let mut probe = StraddleVisitor::default(); + assert!(structural_visit(&root, &mut probe).unwrap().is_none()); + assert_eq!( + probe.events, + vec![ + "enter:node", + "enter:int:1", + "exit:int:1", + "enter:int:2", + "exit:int:2", + "exit:node", + ] + ); +} + +#[derive(Default)] +struct OrderProbe { + events: Vec, +} + +#[dispatch(visit)] +impl OrderProbe { + fn visit_array(&mut self, _array: Array) -> WalkResult { + self.events.push("array".to_string()); + WalkResult::Advance + } + + fn visit_integer(&mut self, value: i64) -> WalkResult { + self.events.push(format!("int:{value}")); + WalkResult::Advance + } +} + +#[test] +fn stateful_structural_walk_supports_post_order() { + let root = Array::new(vec![1i64, 2]); + let mut probe = OrderProbe::default(); + assert!(structural_walk(&root, &mut probe, WalkOrder::PostOrder) + .unwrap() + .is_none()); + assert_eq!(probe.events, vec!["int:1", "int:2", "array"]); +} + +#[test] +fn interrupt_payload_is_returned_to_the_caller() { + let root = Array::new(vec![1i64, 2]); + let outcome = structural_walk( + &root, + |value: &VisitValue| { + if value.cast::() == Some(1) { + return WalkResult::interrupt_with(42i64); + } + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap(); + let Some(interrupt) = outcome else { + panic!("walk unexpectedly completed"); + }; + assert_eq!(i64::try_from(interrupt.value).unwrap(), 42); +} + +#[test] +fn handler_errors_include_native_visit_path() { + let root = Array::new(vec![1i64]); + let error = match structural_walk( + &root, + |value: &VisitValue| -> Result { + if value.cast::().is_some() { + Err(runtime_error("handler failed")) + } else { + Ok(WalkResult::Advance) + } + }, + WalkOrder::PreOrder, + ) { + Err(error) => error, + Ok(_) => panic!("handler unexpectedly succeeded"), + }; + assert_eq!(error.message(), "handler failed"); + assert!(error.backtrace().contains("sequence item [0]")); + assert!(error.backtrace().contains("object `ffi.Array`")); +} + +#[test] +fn visitor_errors_include_native_visit_path() { + struct FailingVisitor; + + impl StructuralVisitor for FailingVisitor { + fn visit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result> { + if value.cast::().is_some() { + return Err(runtime_error("visitor failed")); + } + self.default_visit_children(value, def_region_kind) + } + } + + let root = Array::new(vec![1i64]); + let error = match structural_visit(&root, &mut FailingVisitor) { + Err(error) => error, + Ok(_) => panic!("visitor unexpectedly succeeded"), + }; + assert_eq!(error.message(), "visitor failed"); + assert!(error.backtrace().contains("sequence item [0]")); + assert!(error.backtrace().contains("object `ffi.Array`")); +} + +#[test] +fn visitor_interrupt_propagates_through_default_children() { + struct InterruptingVisitor; + + impl StructuralVisitor for InterruptingVisitor { + fn visit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result> { + if value.cast::() == Some(2) { + return Ok(Some(VisitInterrupt::with(7i64))); + } + self.default_visit_children(value, def_region_kind) + } + } + + let root = Array::new(vec![1i64, 2, 3]); + let outcome = structural_visit(&root, &mut InterruptingVisitor).unwrap(); + let Some(interrupt) = outcome else { + panic!("visitor traversal unexpectedly completed"); + }; + assert_eq!(i64::try_from(interrupt.value).unwrap(), 7); +} + +#[test] +fn closure_walk_receives_def_region_kind() { + // C++: StructuralWalk(root, + // [&](const TVarObj* var, TVMFFIDefRegionKind kind) { ... }) + let root = Array::new(vec![1i64, 2]); + let mut kinds = Vec::new(); + assert!(structural_walk( + &root, + |value: &VisitValue, kind: DefRegionKind| { + if value.cast::().is_some() { + kinds.push(kind); + } + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!(kinds, vec![DefRegionKind::None; 2]); +} + +#[test] +fn closure_walk_supports_post_order_and_skip() { + let root = Array::new(vec![1i64, 2]); + let mut order_probe = Vec::new(); + assert!(structural_walk( + &root, + |value: &VisitValue| { + order_probe.push(value.cast::()); + WalkResult::Advance + }, + WalkOrder::PostOrder, + ) + .unwrap() + .is_none()); + assert_eq!(order_probe, vec![Some(1), Some(2), None]); + + let mut visited = 0; + assert!(structural_walk( + &root, + |value: &VisitValue| { + visited += 1; + if value.cast::().is_none() { + WalkResult::Skip + } else { + WalkResult::Advance + } + }, + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!(visited, 1); +} + +// --------------------------------------------------------------------------- +// Tuple walkers: structural_walk(root, (link1, link2, ...), order) — links +// are tried in order and the first whose argument type matches the value +// runs, the Rust analog of the variadic C++ StructuralWalk callback chain. +// --------------------------------------------------------------------------- +#[test] +fn chain_accepts_owned_object_ref_links() { + let root = Array::new(vec![Array::new(vec![1i64]), Array::new(vec![2i64, 3])]); + let mut lengths = Vec::new(); + assert!(structural_walk( + &root, + ( + |array: Array| { + lengths.push(array.len()); + WalkResult::Advance + }, + |_value: i64| WalkResult::Advance, + ), + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + // The outer Array> fails the strict element check and falls + // through the chain; only the inner arrays match the typed link. + assert_eq!(lengths, vec![1, 2]); +} + +#[test] +fn chain_links_may_mix_def_region_arity() { + // Like #[dispatch(visit)] handlers, each link independently opts into + // the trailing DefRegionKind argument. + let root = Array::new(vec![1i64, 2]); + let mut kinds = Vec::new(); + let mut objects = 0; + assert!(structural_walk( + &root, + ( + |_value: i64, kind: DefRegionKind| { + kinds.push(kind); + WalkResult::Advance + }, + |_value: &VisitValue, kind: DefRegionKind| { + assert_eq!(kind, DefRegionKind::None); + objects += 1; + WalkResult::Advance + }, + ), + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!(kinds, vec![DefRegionKind::None; 2]); + assert_eq!(objects, 1); +} + +#[test] +fn chain_links_can_skip_children() { + let root = Array::new(vec![Array::new(vec![1i64]), Array::new(vec![2i64])]); + let mut arrays = 0; + let mut integers = 0; + assert!(structural_walk( + &root, + ( + |_array: Array| { + arrays += 1; + WalkResult::Skip + }, + |_value: i64| { + integers += 1; + WalkResult::Advance + }, + ), + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!(arrays, 2); + assert_eq!(integers, 0); // both inner arrays were skipped +} + +#[test] +fn chain_link_errors_include_native_visit_path() { + let root = Array::new(vec![1i64]); + let error = match structural_walk( + &root, + ( + |_value: i64| -> Result { Err(runtime_error("link failed")) }, + |_value: &VisitValue| WalkResult::Advance, + ), + WalkOrder::PreOrder, + ) { + Err(error) => error, + Ok(_) => panic!("link unexpectedly succeeded"), + }; + assert_eq!(error.message(), "link failed"); + assert!(error.backtrace().contains("sequence item [0]")); + assert!(error.backtrace().contains("object `ffi.Array`")); +} + +#[test] +fn chain_supports_post_order() { + // Rust borrow rules apply per link: state shared across links goes + // through a RefCell (or a single #[dispatch(visit)] visitor). + let root = Array::new(vec![1i64, 2]); + let events = std::cell::RefCell::new(Vec::new()); + assert!(structural_walk( + &root, + ( + |value: i64| { + events.borrow_mut().push(format!("int:{value}")); + WalkResult::Advance + }, + |_object: &Object| { + events.borrow_mut().push("array".to_string()); + WalkResult::Advance + }, + ), + WalkOrder::PostOrder, + ) + .unwrap() + .is_none()); + assert_eq!(events.into_inner(), vec!["int:1", "int:2", "array"]); +} + +#[derive(Default)] +struct ObjectCounter { + objects: usize, +} + +#[dispatch(visit)] +impl ObjectCounter { + fn visit_object(&mut self, _value: &Object) -> WalkResult { + self.objects += 1; + WalkResult::Advance + } +} + +#[test] +fn chain_splices_dispatch_visitors_between_closures() { + // A `&mut` typed visitor participates in the chain like any other link, + // keeping its own no-match fall-through semantics. + let root = Array::new(vec![1i64, 2]); + let mut counter = ObjectCounter::default(); + let mut integers = 0; + assert!(structural_walk( + &root, + (&mut counter, |_value: i64| { + integers += 1; + WalkResult::Advance + },), + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!(counter.objects, 1); + assert_eq!(integers, 2); +} + +#[test] +fn chain_supports_full_arity() { + // Doubles as the first-match ordering probe: earlier misses fall + // through, the first matching link claims the value, later links + // never run. + let root = Array::new(vec![1i64, 2, 3]); + let mut integers = Vec::new(); + let mut objects = 0; + let mut others = 0; + assert!(structural_walk( + &root, + ( + |_value: f64| WalkResult::Advance, + |_value: bool| WalkResult::Advance, + |_value: tvm_ffi::String| WalkResult::Advance, + |_value: Array| WalkResult::Advance, + |value: i64| { + integers.push(value); + WalkResult::Advance + }, + |_object: &Object, _kind: DefRegionKind| { + objects += 1; + WalkResult::Advance + }, + |_value: &VisitValue, _kind: DefRegionKind| { + others += 1; + WalkResult::Advance + }, + |_value: &VisitValue| WalkResult::Advance, + ), + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!(integers, vec![1, 2, 3]); + assert_eq!(objects, 1); // the array itself; integers matched earlier + assert_eq!(others, 0); // every value matched an earlier link +} + +#[test] +fn typed_lambda_walks_bare_and_as_single_link_tuple() { + // A lone typed handler needs no tuple: unmatched values (the array + // itself) advance normally. The 1-tuple spelling routes through the + // chain impls instead and must agree. + let root = Array::new(vec![1i64, 2, 3]); + let mut bare = 0; + assert!(structural_walk( + &root, + |value: i64| { + bare += value; + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + let mut tupled = 0; + assert!(structural_walk( + &root, + (|value: i64| { + tupled += value; + WalkResult::Advance + },), + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!((bare, tupled), (6, 6)); +} + +#[test] +fn bare_node_lambda_takes_def_region_kind() { + let root = Array::new(vec![1i64, 2]); + let mut objects = 0; + assert!(structural_walk( + &root, + |_object: &Object, kind: DefRegionKind| { + assert_eq!(kind, DefRegionKind::None); + objects += 1; + WalkResult::Advance + }, + WalkOrder::PreOrder, + ) + .unwrap() + .is_none()); + assert_eq!(objects, 1); +} diff --git a/rust/tvm-ffi/tests/test_structural_visitor_alignment.rs b/rust/tvm-ffi/tests/test_structural_visitor_alignment.rs new file mode 100644 index 000000000..2d964bdff --- /dev/null +++ b/rust/tvm-ffi/tests/test_structural_visitor_alignment.rs @@ -0,0 +1,133 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +//! Rust mirror of the C++ visitor example. +//! +//! `RecordingVisitor` lines up member-for-member with the C++ +//! `TestVisitorObj` (tests/cpp/extra/test_structural_visit.cc): its `visit` +//! plays the role of the C++ vtable `VisitImpl`, and the array arm plays the +//! role of the C++ `TFuncObj::StructuralVisit` hook +//! (tests/cpp/testing_object.h): the first element is visited as a recursive +//! definition region, the rest inherit the surrounding state. + +use tvm_ffi::{ + structural_visit, Array, DefRegionKind, Result, String as FfiString, StructuralVisitor, + VisitInterrupt, VisitValue, +}; + +/// C++: class TestVisitorObj : public StructuralVisitorObj +#[derive(Default)] +struct RecordingVisitor { + /// C++: `std::vector visited;` + visited: Vec, + /// C++: `std::vector modes;` + modes: Vec, + /// C++: `ObjectRef interrupt_on;` + interrupt_on: Option, +} + +impl StructuralVisitor for RecordingVisitor { + /// C++ analog: `TestVisitorObj::VisitImpl` — record every value together + /// with the active def-region state, optionally interrupt with a payload, + /// otherwise delegate recursion explicitly. The array arm mirrors + /// `TFuncObj::StructuralVisit`. + fn visit( + &mut self, + value: &VisitValue, + def_region_kind: DefRegionKind, + ) -> Result> { + let integer = value.cast::(); + let label = match integer { + Some(integer) => integer.to_string(), + None if value.cast::>().is_some() => "array".to_string(), + None => "obj".to_string(), + }; + // C++: visited.push_back(value_ref); + // modes.push_back(def_region_mode_); + self.visited.push(label); + self.modes.push(def_region_kind); + + // C++: if (value_ref.same_as(interrupt_on)) + // return VisitInterrupt(String("stop")); + if self.interrupt_on.is_some() && integer == self.interrupt_on { + return Ok(Some(VisitInterrupt::with(FfiString::from("stop")))); + } + + // C++ analog: `TFuncObj::StructuralVisit` — visit "params" + // (element 0) under a recursive definition region, then the "body" + // (element 1) under the inherited state. + if let Some(array) = value.cast::>() { + // C++: visitor->WithDefRegionKind(kTVMFFIDefRegionKindRecursive, + // [&] { return visitor->VisitExpected(self->params); }) + let params = array.get(0).unwrap(); + if let Some(interrupt) = self.visit_child(¶ms, DefRegionKind::Recursive)? { + return Ok(Some(interrupt)); + } + // C++: visitor->VisitExpected(self->body) (inherits the state) + let body = array.get(1).unwrap(); + return self.visit_child(&body, def_region_kind); + } + + // C++: return DefaultVisitExpected(value); + self.default_visit_children(value, def_region_kind) + } +} + +/// C++ analog: TEST(StructuralVisitor, TraversesFunction) — the def-region +/// state flips to Recursive under "params" and back to None for the "body". +#[test] +fn records_values_and_def_region_modes() { + let root = Array::new(vec![10i64, 20]); + let mut visitor = RecordingVisitor::default(); + + let outcome = structural_visit(&root, &mut visitor).unwrap(); + + assert!(outcome.is_none()); + assert_eq!(visitor.visited, vec!["array", "10", "20"]); + assert_eq!( + visitor.modes, + vec![ + DefRegionKind::None, // the array itself + DefRegionKind::Recursive, // element 0: the "params" position + DefRegionKind::None, // element 1: the "body" position + ] + ); +} + +/// C++ analog: TEST(StructuralVisitor, StopsOnInterrupt) — the traversal +/// halts at the marked value and the payload reaches the caller. +#[test] +fn stops_on_interrupt_with_payload() { + let root = Array::new(vec![10i64, 20]); + let mut visitor = RecordingVisitor { + interrupt_on: Some(20), + ..RecordingVisitor::default() + }; + + let outcome = structural_visit(&root, &mut visitor).unwrap(); + + let Some(interrupt) = outcome else { + panic!("traversal unexpectedly completed"); + }; + assert_eq!( + FfiString::try_from(interrupt.value).unwrap().as_str(), + "stop" + ); + assert_eq!(visitor.visited, vec!["array", "10", "20"]); +}