This is an automated email from the ASF dual-hosted git repository.
tqchen pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm-ffi.git
The following commit(s) were added to refs/heads/main by this push:
new 556514c7 [Test][Rust] Use existing C++ registrations for structural
walk and map (#745)
556514c7 is described below
commit 556514c76225cb35a55e0649098b07be3fba2991
Author: Shushi Hong <[email protected]>
AuthorDate: Sat Sep 5 08:33:10 2026 -0400
[Test][Rust] Use existing C++ registrations for structural walk and map
(#745)
Rust walk/map tests currently register custom types and hooks while the
test harness runs tests concurrently. Registering these hooks can
relocate the shared type-attribute columns read by other tests; #743
addressed the visit-test race with a per-test initialization barrier.
Use the existing C++ container hooks and the startup-registered
`testing.TestObjectBase` and `testing.TestNonCopyable` types instead.
Remove the Rust-only test objects, type/field/hook registration helpers,
and initialization barriers from both structural test files. Treat
`DLDataType` as a primitive leaf without registering a traversal hook.
Keep tests in their original integration-test files, with no C++ or Rust
implementation changes.
This narrows the test scope to traversal through existing C++
registrations. Ordinary walk/map dispatch, recursion, copy-on-write,
interrupts, errors, and panic recovery remain covered. Reflection checks
now use the existing C++ objects. Remove 16 specialized tests that
depend on custom hook handles, DAG/FreeVar metadata, or injected
getter/setter failures; the rewritten tests also stop asserting special
field flags and exact shallow-copy callback counts. These checks are
removed, not relocated or claimed to be equivalently covered.
Validation on Linux x86_64:
- Rust workspace: 238 tests passed, including 45 visit/walk and 35
mutate/map tests; no compiler warnings.
- Both structural test suites passed 100 concurrent rounds, each
executable using 8 test threads.
- Neither executable imports `TVMFFITypeRegister*` or
`TVMFFITypeGetOrAllocIndex`; both link the existing C++ testing library.
- Rust formatting, ASF headers, file-type checks, and `git diff --check`
passed.
Windows/MSVC validation remains pending CI for this change.
---
rust/tvm-ffi/tests/test_structural_mutate.rs | 1050 ++------------------------
rust/tvm-ffi/tests/test_structural_visit.rs | 727 ++----------------
2 files changed, 151 insertions(+), 1626 deletions(-)
diff --git a/rust/tvm-ffi/tests/test_structural_mutate.rs
b/rust/tvm-ffi/tests/test_structural_mutate.rs
index f8c4af8a..feadbebb 100644
--- a/rust/tvm-ffi/tests/test_structural_mutate.rs
+++ b/rust/tvm-ffi/tests/test_structural_mutate.rs
@@ -18,518 +18,16 @@
*/
use std::cell::{Cell, RefCell};
-use std::sync::atomic::{AtomicUsize, Ordering};
-use std::sync::{LazyLock, Mutex};
-
-use tvm_ffi::derive::{Object as DeriveObject, ObjectRef as DeriveObjectRef};
+use tvm_ffi::collections::map::MapObj;
+use tvm_ffi::function::FunctionObj;
use tvm_ffi::object::ObjectRef;
-use tvm_ffi::tvm_ffi_sys::{
- TVMFFIAny, TVMFFIAnyViewToOwnedAny, TVMFFIByteArray,
TVMFFIFieldFlagBitMask, TVMFFIFieldInfo,
- TVMFFISEqHashKind, TVMFFITypeMetadata, TVMFFITypeRegisterAttr,
-};
use tvm_ffi::{
dispatch, structural_map, structural_mutate, Any, AnyView, Array,
CallbackMutator,
- DefRegionKind, Error, Function, InplaceValue, Map, MapDispatch, MapValue,
MutateCallbacks,
- Mutator, Object, ObjectArc, ObjectCore, ObjectRefCore, Result, String as
FfiString,
+ DefRegionKind, Error, FieldGetter, Function, InplaceValue, Map,
MapDispatch, MapValue,
+ MutateCallbacks, Mutator, Object, ObjectArc, ObjectRefCore, Result, String
as FfiString,
StructuralMutator, StructuralVarRemap, TypeIndex, WalkOrder, RUNTIME_ERROR,
};
-// These registration entry points are needed only to build reflected test
-// types. Keep them local instead of expanding tvm-ffi-sys's public API.
-unsafe extern "C" {
- fn TVMFFITypeGetOrAllocIndex(
- type_key: *const TVMFFIByteArray,
- static_type_index: i32,
- type_depth: i32,
- num_child_slots: i32,
- child_slots_can_overflow: i32,
- parent_type_index: i32,
- ) -> i32;
- fn TVMFFITypeRegisterField(type_index: i32, info: *const TVMFFIFieldInfo)
-> i32;
- fn TVMFFITypeRegisterMetadata(type_index: i32, metadata: *const
TVMFFITypeMetadata) -> i32;
-}
-
-#[repr(C)]
-#[derive(DeriveObject)]
-#[type_key = "testing.RustStructuralDagNode"]
-#[type_final]
-struct RustDagNodeObj {
- base: Object,
-}
-
-#[repr(C)]
-#[derive(DeriveObjectRef, Clone)]
-struct RustDagNode {
- data: ObjectArc<RustDagNodeObj>,
-}
-
-#[repr(C)]
-#[derive(DeriveObject)]
-#[type_key = "testing.RustStructuralFreeVar"]
-#[type_final]
-struct RustFreeVarObj {
- base: Object,
-}
-
-#[repr(C)]
-#[derive(DeriveObjectRef, Clone)]
-struct RustFreeVar {
- data: ObjectArc<RustFreeVarObj>,
-}
-
-#[repr(C)]
-#[derive(DeriveObject)]
-#[type_key = "testing.RustStructuralPair"]
-#[type_final]
-struct RustPairObj {
- base: Object,
- first: Any,
- ignored: Any,
-}
-
-#[repr(C)]
-#[derive(DeriveObjectRef, Clone)]
-struct RustPair {
- data: ObjectArc<RustPairObj>,
-}
-
-#[repr(C)]
-#[derive(DeriveObject)]
-#[type_key = "testing.RustStructuralNoCopy"]
-#[type_final]
-struct RustNoCopyObj {
- base: Object,
-}
-
-#[repr(C)]
-#[derive(DeriveObjectRef, Clone)]
-struct RustNoCopy {
- data: ObjectArc<RustNoCopyObj>,
-}
-
-#[repr(C)]
-#[derive(DeriveObject)]
-#[type_key = "testing.RustStructuralHookNode"]
-#[type_final]
-struct RustHookNodeObj {
- base: Object,
-}
-
-#[repr(C)]
-#[derive(DeriveObjectRef, Clone)]
-struct RustHookNode {
- data: ObjectArc<RustHookNodeObj>,
-}
-
-#[repr(C)]
-#[derive(DeriveObject)]
-#[type_key = "testing.RustStructuralFailingGetter"]
-#[type_final]
-struct RustFailingGetterObj {
- base: Object,
- value: Any,
-}
-
-#[repr(C)]
-#[derive(DeriveObjectRef, Clone)]
-struct RustFailingGetter {
- data: ObjectArc<RustFailingGetterObj>,
-}
-
-#[repr(C)]
-#[derive(DeriveObject)]
-#[type_key = "testing.RustStructuralFailingSetter"]
-#[type_final]
-struct RustFailingSetterObj {
- base: Object,
- value: Any,
-}
-
-#[repr(C)]
-#[derive(DeriveObjectRef, Clone)]
-struct RustFailingSetter {
- data: ObjectArc<RustFailingSetterObj>,
-}
-
-static SHALLOW_COPY_CALLS: AtomicUsize = AtomicUsize::new(0);
-static REGISTERED_MUTATE_CALLS: AtomicUsize = AtomicUsize::new(0);
-static REGISTERED_MAYBE_INPLACE_MUTATE_CALLS: AtomicUsize =
AtomicUsize::new(0);
-static REFLECTED_TEST_LOCK: Mutex<()> = Mutex::new(());
-static REGISTERED_HOOK_TEST_LOCK: Mutex<()> = Mutex::new(());
-
-thread_local! {
- static RETAINED_MUTATOR: RefCell<Option<Any>> = const { RefCell::new(None)
};
- static PROBE_FOREIGN_THREAD_MUTATOR: Cell<bool> = const { Cell::new(false)
};
- static FOREIGN_THREAD_MUTATOR_ERROR: RefCell<Option<String>> = const {
RefCell::new(None) };
-}
-
-fn call_mutator_from_foreign_thread(mutator: AnyView<'_>) -> String {
- // Keep an owning reference on this thread while the worker constructs a
- // borrowed ABI view from the raw object address.
- let mut owner = Any::from(mutator);
- let raw = unsafe { *Any::as_data_ptr(&mut owner) };
- let type_index = raw.type_index;
- let object = unsafe { raw.data_union.v_obj } as usize;
- std::thread::spawn(move || {
- let mut raw = TVMFFIAny::new();
- raw.type_index = type_index;
- raw.data_union.v_obj = object as *mut _;
- let borrowed = std::mem::ManuallyDrop::new(unsafe {
Any::from_raw_ffi_any(raw) });
- match Function::get_global("ffi.StructuralMutatorMutate")
- .unwrap()
- .call_packed(&[AnyView::from(&*borrowed), AnyView::from(&1i64)])
- {
- Err(error) => error.message().to_string(),
- Ok(_) => "foreign-thread mutator call unexpectedly
succeeded".to_string(),
- }
- })
- .join()
- .unwrap()
-}
-
-fn run_registered_mutate_hook(args: &[AnyView<'_>], calls: &AtomicUsize) ->
Result<Any> {
- assert_eq!(args.len(), 2);
- if PROBE_FOREIGN_THREAD_MUTATOR.with(Cell::get) {
- let message = call_mutator_from_foreign_thread(args[0]);
- FOREIGN_THREAD_MUTATOR_ERROR.with(|error|
error.replace(Some(message)));
- return Ok(Any::from(args[1]));
- }
- Function::get_global("ffi.StructuralMutatorDefRegionKind")
- .unwrap()
- .call_packed(&[args[0]])?;
-
- // Keep one reference so the test can verify that a type hook cannot use
- // the mutator after the structural-mutation call has ended.
- RETAINED_MUTATOR.with(|retained| {
- retained.replace(Some(Any::from(args[0])));
- });
-
- let cached = Function::get_global("ffi.StructuralMutatorVarRemapGet")
- .unwrap()
- .call_packed(&[args[0], args[1]])?;
- if cached.type_index() != TypeIndex::kTVMFFINone as i32 {
- return Ok(cached);
- }
-
- let child = Any::from(1i64);
- let mutated = Function::get_global("ffi.StructuralMutatorMutate")
- .unwrap()
- .call_packed(&[args[0], AnyView::from(&child)])?;
- Function::get_global("ffi.StructuralMutatorVarRemapSet")
- .unwrap()
- .call_packed(&[args[0], args[1], AnyView::from(&mutated)])?;
- calls.fetch_add(1, Ordering::Relaxed);
- Ok(mutated)
-}
-
-unsafe extern "C" fn any_field_getter(field: *mut std::ffi::c_void, result:
*mut TVMFFIAny) -> i32 {
- TVMFFIAnyViewToOwnedAny(field.cast(), result)
-}
-
-unsafe extern "C" fn any_field_setter(
- field: *mut std::ffi::c_void,
- value: *const TVMFFIAny,
-) -> i32 {
- let mut replacement = TVMFFIAny::new();
- let code = TVMFFIAnyViewToOwnedAny(value, &mut replacement);
- if code != 0 {
- return code;
- }
- let field = &mut *field.cast::<Any>();
- *field = Any::from_raw_ffi_any(replacement);
- 0
-}
-
-unsafe extern "C" fn clone_any_then_fail(
- source: *mut std::ffi::c_void,
- result: *mut TVMFFIAny,
-) -> i32 {
- let code = TVMFFIAnyViewToOwnedAny(source.cast(), result);
- if code != 0 {
- return code;
- }
- Error::set_raised(&Error::new(
- RUNTIME_ERROR,
- "callback failed after writing an owning result",
- "",
- ));
- -1
-}
-
-unsafe extern "C" fn setter_safe_call(
- _handle: *mut std::ffi::c_void,
- args: *const TVMFFIAny,
- _num_args: i32,
- result: *mut TVMFFIAny,
-) -> i32 {
- clone_any_then_fail(args.add(1).cast_mut().cast(), result)
-}
-
-fn register_any_field(type_index: i32, name: &'static str, offset: usize,
flags: i64) {
- let field = TVMFFIFieldInfo {
- name: unsafe { TVMFFIByteArray::from_str(name) },
- doc: unsafe { TVMFFIByteArray::from_str("Rust structural-mutation test
field") },
- metadata: unsafe { TVMFFIByteArray::from_str("") },
- flags,
- size: std::mem::size_of::<Any>() as i64,
- alignment: std::mem::align_of::<Any>() as i64,
- offset: offset as i64,
- getter: Some(any_field_getter),
- setter: any_field_setter as *mut std::ffi::c_void,
- default_value_or_factory: TVMFFIAny::new(),
- field_static_type_index: -1,
- };
- assert_eq!(unsafe { TVMFFITypeRegisterField(type_index, &field) }, 0);
-}
-
-fn register_test_type(type_key: &'static str, total_size: usize, kind:
TVMFFISEqHashKind) -> i32 {
- let type_key = unsafe { TVMFFIByteArray::from_str(type_key) };
- let type_index = unsafe {
- TVMFFITypeGetOrAllocIndex(
- &type_key,
- -1,
- Object::TYPE_DEPTH + 1,
- 0,
- 1,
- Object::type_index(),
- )
- };
- assert!(type_index >= TypeIndex::kTVMFFIDynObjectBegin as i32);
- let metadata = TVMFFITypeMetadata {
- doc: unsafe { TVMFFIByteArray::from_str("Rust structural-mutation test
type") },
- creator: None,
- total_size: i32::try_from(total_size).unwrap(),
- structural_eq_hash_kind: kind as i32,
- };
- assert_eq!(
- unsafe { TVMFFITypeRegisterMetadata(type_index, &metadata) },
- 0
- );
- type_index
-}
-
-fn register_function_attr(type_index: i32, name: &'static str, function:
Function) {
- let name = unsafe { TVMFFIByteArray::from_str(name) };
- let mut value = Any::from(function);
- assert_eq!(
- unsafe { TVMFFITypeRegisterAttr(type_index, &name,
Any::as_data_ptr(&mut value)) },
- 0
- );
-}
-
-static REGISTER_TEST_TYPES: LazyLock<()> = LazyLock::new(|| {
- register_test_type(
- RustDagNodeObj::TYPE_KEY,
- std::mem::size_of::<RustDagNodeObj>(),
- TVMFFISEqHashKind::kTVMFFISEqHashKindDAGNode,
- );
- register_test_type(
- RustFreeVarObj::TYPE_KEY,
- std::mem::size_of::<RustFreeVarObj>(),
- TVMFFISEqHashKind::kTVMFFISEqHashKindFreeVar,
- );
- register_test_type(
- RustNoCopyObj::TYPE_KEY,
- std::mem::size_of::<RustNoCopyObj>(),
- TVMFFISEqHashKind::kTVMFFISEqHashKindTreeNode,
- );
- register_test_type(
- RustHookNodeObj::TYPE_KEY,
- std::mem::size_of::<RustHookNodeObj>(),
- TVMFFISEqHashKind::kTVMFFISEqHashKindDAGNode,
- );
-
- let pair_type_index = register_test_type(
- RustPairObj::TYPE_KEY,
- std::mem::size_of::<RustPairObj>(),
- TVMFFISEqHashKind::kTVMFFISEqHashKindTreeNode,
- );
- register_any_field(
- pair_type_index,
- "first",
- std::mem::offset_of!(RustPairObj, first),
- TVMFFIFieldFlagBitMask::kTVMFFIFieldFlagBitMaskSEqHashDefRecursive as
i64,
- );
- register_any_field(
- pair_type_index,
- "ignored",
- std::mem::offset_of!(RustPairObj, ignored),
- TVMFFIFieldFlagBitMask::kTVMFFIFieldFlagBitMaskSEqHashIgnore as i64,
- );
-
- let getter_type_index = register_test_type(
- RustFailingGetterObj::TYPE_KEY,
- std::mem::size_of::<RustFailingGetterObj>(),
- TVMFFISEqHashKind::kTVMFFISEqHashKindTreeNode,
- );
- let field = TVMFFIFieldInfo {
- name: unsafe { TVMFFIByteArray::from_str("value") },
- doc: unsafe { TVMFFIByteArray::from_str("Fail after producing an
owning field value") },
- metadata: unsafe { TVMFFIByteArray::from_str("") },
- flags: 0,
- size: std::mem::size_of::<Any>() as i64,
- alignment: std::mem::align_of::<Any>() as i64,
- offset: std::mem::offset_of!(RustFailingGetterObj, value) as i64,
- getter: Some(clone_any_then_fail),
- setter: any_field_setter as *mut std::ffi::c_void,
- default_value_or_factory: TVMFFIAny::new(),
- field_static_type_index: -1,
- };
- assert_eq!(
- unsafe { TVMFFITypeRegisterField(getter_type_index, &field) },
- 0
- );
- let shallow_copy = Function::from_packed(|args| {
- let source = RustFailingGetter::try_from(args[0])?;
- Ok(Any::from(RustFailingGetter {
- data: ObjectArc::new(RustFailingGetterObj {
- base: Object::new(),
- value: source.data.value.clone(),
- }),
- }))
- });
- register_function_attr(getter_type_index, "__ffi_shallow_copy__",
shallow_copy);
-
- let setter_type_index = register_test_type(
- RustFailingSetterObj::TYPE_KEY,
- std::mem::size_of::<RustFailingSetterObj>(),
- TVMFFISEqHashKind::kTVMFFISEqHashKindTreeNode,
- );
- let setter = unsafe { Function::from_extern_c(std::ptr::null_mut(),
setter_safe_call, None) };
- let field = TVMFFIFieldInfo {
- name: unsafe { TVMFFIByteArray::from_str("value") },
- doc: unsafe { TVMFFIByteArray::from_str("Fail after producing an
owning setter result") },
- metadata: unsafe { TVMFFIByteArray::from_str("") },
- flags: TVMFFIFieldFlagBitMask::kTVMFFIFieldFlagBitSetterIsFunctionObj
as i64,
- size: std::mem::size_of::<Any>() as i64,
- alignment: std::mem::align_of::<Any>() as i64,
- offset: std::mem::offset_of!(RustFailingSetterObj, value) as i64,
- getter: Some(any_field_getter),
- setter: unsafe {
- ObjectArc::as_raw(<Function as ObjectRefCore>::data(&setter))
- .cast_mut()
- .cast()
- },
- default_value_or_factory: TVMFFIAny::new(),
- field_static_type_index: -1,
- };
- assert_eq!(
- unsafe { TVMFFITypeRegisterField(setter_type_index, &field) },
- 0
- );
- let shallow_copy = Function::from_packed(|args| {
- let source = RustFailingSetter::try_from(args[0])?;
- Ok(Any::from(RustFailingSetter {
- data: ObjectArc::new(RustFailingSetterObj {
- base: Object::new(),
- value: source.data.value.clone(),
- }),
- }))
- });
- register_function_attr(setter_type_index, "__ffi_shallow_copy__",
shallow_copy);
-
- let shallow_copy = Function::from_packed(|args| {
- SHALLOW_COPY_CALLS.fetch_add(1, Ordering::Relaxed);
- let source = RustPair::try_from(args[0])?;
- Ok(Any::from(RustPair {
- data: ObjectArc::new(RustPairObj {
- base: Object::new(),
- first: source.data.first.clone(),
- ignored: source.data.ignored.clone(),
- }),
- }))
- });
- register_function_attr(pair_type_index, "__ffi_shallow_copy__",
shallow_copy);
-
- let registered_mutate =
- Function::from_packed(|args| run_registered_mutate_hook(args,
®ISTERED_MUTATE_CALLS));
- register_function_attr(
- RustHookNodeObj::type_index(),
- "__s_mutate__",
- registered_mutate,
- );
-
- let registered_maybe_inplace_mutate = Function::from_packed(|args| {
- run_registered_mutate_hook(args,
®ISTERED_MAYBE_INPLACE_MUTATE_CALLS)
- });
- register_function_attr(
- RustHookNodeObj::type_index(),
- "__s_maybe_inplace_mutate__",
- registered_maybe_inplace_mutate,
- );
-});
-
-fn ensure_test_types_registered() {
- LazyLock::force(®ISTER_TEST_TYPES);
-}
-
-fn rust_dag_node() -> RustDagNode {
- LazyLock::force(®ISTER_TEST_TYPES);
- RustDagNode {
- data: ObjectArc::new(RustDagNodeObj {
- base: Object::new(),
- }),
- }
-}
-
-fn rust_free_var() -> RustFreeVar {
- LazyLock::force(®ISTER_TEST_TYPES);
- RustFreeVar {
- data: ObjectArc::new(RustFreeVarObj {
- base: Object::new(),
- }),
- }
-}
-
-fn rust_pair(first: impl Into<Any>, ignored: impl Into<Any>) -> RustPair {
- LazyLock::force(®ISTER_TEST_TYPES);
- RustPair {
- data: ObjectArc::new(RustPairObj {
- base: Object::new(),
- first: first.into(),
- ignored: ignored.into(),
- }),
- }
-}
-
-fn rust_no_copy() -> RustNoCopy {
- ensure_test_types_registered();
- RustNoCopy {
- data: ObjectArc::new(RustNoCopyObj {
- base: Object::new(),
- }),
- }
-}
-
-fn rust_hook_node() -> RustHookNode {
- ensure_test_types_registered();
- RustHookNode {
- data: ObjectArc::new(RustHookNodeObj {
- base: Object::new(),
- }),
- }
-}
-
-fn rust_failing_getter(value: impl Into<Any>) -> RustFailingGetter {
- ensure_test_types_registered();
- RustFailingGetter {
- data: ObjectArc::new(RustFailingGetterObj {
- base: Object::new(),
- value: value.into(),
- }),
- }
-}
-
-fn rust_failing_setter(value: impl Into<Any>) -> RustFailingSetter {
- ensure_test_types_registered();
- RustFailingSetter {
- data: ObjectArc::new(RustFailingSetterObj {
- base: Object::new(),
- value: value.into(),
- }),
- }
-}
-
struct IncrementIntegers;
impl MapDispatch for IncrementIntegers {
@@ -608,44 +106,6 @@ impl StructuralMutator for ReplaceNone {
}
}
-struct RemappingFreeVar {
- remap: StructuralVarRemap,
- type_index: i32,
- calls: usize,
-}
-
-impl StructuralMutator for RemappingFreeVar {
- fn dispatch_mutate(&mut self, value: &MapValue, def_region_kind:
DefRegionKind) -> Result<Any> {
- if value.type_index() == self.type_index {
- if let Some(mutated) = self.remap.get(value)? {
- return Ok(mutated);
- }
- self.calls += 1;
- let mutated = Any::from(41i64);
- self.remap.set(value, &mutated)?;
- Ok(mutated)
- } else {
- self.default_mutate(value, def_region_kind)
- }
- }
-
- fn dispatch_maybe_inplace_mutate(
- &mut self,
- value: InplaceValue<'_>,
- def_region_kind: DefRegionKind,
- ) -> Result<Any> {
- self.default_maybe_inplace_mutate(value, def_region_kind)
- }
-
- fn var_remap_get(&mut self, var: &MapValue) -> Result<Option<Any>> {
- self.remap.get(var)
- }
-
- fn var_remap_set(&mut self, var: &MapValue, mutated_value: &Any) ->
Result<()> {
- self.remap.set(var, mutated_value)
- }
-}
-
struct RecursiveEntryMutator {
remap: StructuralVarRemap,
use_owned_value: bool,
@@ -686,48 +146,32 @@ impl StructuralMutator for RecursiveEntryMutator {
}
}
-#[derive(Default)]
-struct RejectAliasedReentry {
- remap: StructuralVarRemap,
- rejected: bool,
+fn reflected_object() -> Any {
+ // Reference the existing test library so its C++ startup registrations
are linked.
+ assert_eq!(
+ unsafe { tvm_ffi::tvm_ffi_sys::TVMFFITestingDummyTarget() },
+ 0
+ );
+ Function::get_global("ffi.MakeObjectFromPackedArgs")
+ .unwrap()
+ .call_tuple((
+ FfiString::from("testing.TestObjectBase"),
+ FfiString::from("v_i64"),
+ 1i64,
+ FfiString::from("v_f64"),
+ 2.5f64,
+ FfiString::from("v_str"),
+ FfiString::from("a reflected string"),
+ ))
+ .unwrap()
}
-impl StructuralMutator for RejectAliasedReentry {
- fn dispatch_mutate(&mut self, value: &MapValue, def_region_kind:
DefRegionKind) -> Result<Any> {
- if let Some(integer) = value.cast::<i64>() {
- let retained = RETAINED_MUTATOR.with(|slot|
slot.borrow().as_ref().unwrap().clone());
- let error = match
Function::get_global("ffi.StructuralMutatorMutate")
- .unwrap()
- .call_packed(&[AnyView::from(&retained),
AnyView::from(&integer)])
- {
- Ok(_) => panic!("aliased structural-mutator reentry
unexpectedly succeeded"),
- Err(error) => error,
- };
- assert!(error
- .message()
- .contains("may only be called by its active registered hook"));
- self.rejected = true;
- Ok(Any::from(integer + 1))
- } else {
- self.default_mutate(value, def_region_kind)
- }
- }
-
- fn dispatch_maybe_inplace_mutate(
- &mut self,
- value: InplaceValue<'_>,
- def_region_kind: DefRegionKind,
- ) -> Result<Any> {
- self.default_maybe_inplace_mutate(value, def_region_kind)
- }
-
- fn var_remap_get(&mut self, var: &MapValue) -> Result<Option<Any>> {
- self.remap.get(var)
- }
-
- fn var_remap_set(&mut self, var: &MapValue, mutated_value: &Any) ->
Result<()> {
- self.remap.set(var, mutated_value)
- }
+fn reflected_field<T: TryFrom<Any, Error = Error>>(value: &Any, name: &str) ->
T {
+ let object = ObjectRef::try_from(value.clone()).unwrap();
+ FieldGetter::new(value.type_index(), name)
+ .unwrap()
+ .get::<_, T>(&**ObjectRef::data(&object))
+ .unwrap()
}
fn array_pointer<T>(array: &Array<T>) -> *const
tvm_ffi::collections::array::ArrayObj
@@ -779,7 +223,6 @@ fn dict_item(dict: &Any, key: i64) -> i64 {
#[test]
fn unique_array_is_reused_while_shared_array_uses_copy_on_write() {
- ensure_test_types_registered();
let unique = Array::new(vec![1i64, 2, 3]);
let unique_pointer = array_pointer(&unique);
let mapped = structural_map(unique, &mut IncrementIntegers,
WalkOrder::PostOrder)
@@ -800,7 +243,6 @@ fn
unique_array_is_reused_while_shared_array_uses_copy_on_write() {
#[test]
fn user_driven_mutator_controls_default_recursion_and_in_place_opt_in() {
- ensure_test_types_registered();
let unique = Array::new(vec![1i64, 2]);
let unique_pointer = array_pointer(&unique);
let mutated =
@@ -859,24 +301,6 @@ fn
none_values_are_dispatched_to_map_callbacks_and_user_mutators() {
assert_eq!(mutator.calls, 1);
}
-#[test]
-fn user_mutator_can_store_a_changed_free_var_result() {
- ensure_test_types_registered();
- let var = rust_free_var();
- let type_index = RustFreeVarObj::type_index();
- let root = call_global("ffi.Array", &[Any::from(var.clone()),
Any::from(var)]);
- let mut mutator = RemappingFreeVar {
- remap: StructuralVarRemap::default(),
- type_index,
- calls: 0,
- };
-
- let mutated = structural_mutate(root, &mut mutator).unwrap();
- assert_eq!(mutator.calls, 1);
- assert_eq!(i64::try_from(array_item(&mutated, 0)).unwrap(), 41);
- assert_eq!(i64::try_from(array_item(&mutated, 1)).unwrap(), 41);
-}
-
#[test]
fn user_mutator_recursive_entries_reenter_the_same_mutator() {
let mut borrowed = RecursiveEntryMutator {
@@ -905,69 +329,8 @@ fn
user_mutator_recursive_entries_reenter_the_same_mutator() {
}
#[test]
-fn dag_identity_caches_the_final_pre_order_replacement() {
- ensure_test_types_registered();
- let node = rust_dag_node();
- let root = call_global("ffi.Array", &[Any::from(node.clone()),
Any::from(node)]);
- let mut identity_calls = 0;
- let mapped = structural_map(
- root,
- (
- |_node: &RustDagNodeObj| {
- identity_calls += 1;
- Any::from(Array::new(vec![1i64]))
- },
- |integer: i64| Any::from(integer + 1),
- ),
- WalkOrder::PreOrder,
- )
- .unwrap();
-
- let first = array_item(&mapped, 0);
- let second = array_item(&mapped, 1);
- assert_eq!(identity_calls, 1);
- assert_eq!(any_object_pointer(&first), any_object_pointer(&second));
- assert_eq!(i64::try_from(array_item(&first, 0)).unwrap(), 2);
- assert_eq!(i64::try_from(array_item(&second, 0)).unwrap(), 2);
-}
-
-#[test]
-fn free_var_identity_caches_the_final_pre_order_replacement() {
- ensure_test_types_registered();
- let node = rust_free_var();
- let type_index = RustFreeVarObj::type_index();
- let root = call_global("ffi.Array", &[Any::from(node.clone()),
Any::from(node)]);
- let mut identity_calls = 0;
- let mapped = structural_map(
- root,
- |value: &MapValue| {
- if value.type_index() == type_index {
- identity_calls += 1;
- Any::from(Array::new(vec![1i64]))
- } else if let Some(integer) = value.cast::<i64>() {
- Any::from(integer + 1)
- } else {
- value.to_owned()
- }
- },
- WalkOrder::PreOrder,
- )
- .unwrap();
-
- let first = array_item(&mapped, 0);
- let second = array_item(&mapped, 1);
- assert_eq!(identity_calls, 1);
- assert_eq!(any_object_pointer(&first), any_object_pointer(&second));
- assert_eq!(i64::try_from(array_item(&first, 0)).unwrap(), 2);
- assert_eq!(i64::try_from(array_item(&second, 0)).unwrap(), 2);
-}
-
-#[test]
-fn reflected_fields_use_shallow_copy_setters_and_field_flags() {
- ensure_test_types_registered();
- let _guard = REFLECTED_TEST_LOCK.lock().unwrap();
- let source = rust_pair(1i64, 9i64);
- let source_pointer = unsafe { ObjectArc::as_raw(&source.data) };
+fn reflected_fields_use_shallow_copy_and_setters() {
+ let source = reflected_object();
let mut regions = Vec::new();
let mapped = structural_map(
source.clone(),
@@ -977,99 +340,58 @@ fn
reflected_fields_use_shallow_copy_setters_and_field_flags() {
},
WalkOrder::PostOrder,
)
- .and_then(RustPair::try_from)
.unwrap();
- assert_ne!(unsafe { ObjectArc::as_raw(&mapped.data) }, source_pointer);
- assert_eq!(i64::try_from(source.data.first.clone()).unwrap(), 1);
- assert_eq!(i64::try_from(mapped.data.first.clone()).unwrap(), 2);
- assert_eq!(i64::try_from(mapped.data.ignored.clone()).unwrap(), 9);
- assert_eq!(regions, vec![DefRegionKind::Recursive]);
+ assert_ne!(any_object_pointer(&mapped), any_object_pointer(&source));
+ assert_eq!(reflected_field::<i64>(&source, "v_i64"), 1);
+ assert_eq!(reflected_field::<i64>(&mapped, "v_i64"), 2);
+ assert_eq!(reflected_field::<f64>(&mapped, "v_f64"), 2.5);
+ assert_eq!(
+ reflected_field::<FfiString>(&mapped, "v_str").as_str(),
+ "a reflected string"
+ );
+ assert_eq!(regions, vec![DefRegionKind::None]);
}
#[test]
-fn reflected_no_change_still_validates_copy_and_returns_original() {
- ensure_test_types_registered();
- let _guard = REFLECTED_TEST_LOCK.lock().unwrap();
- let source = rust_pair(1i64, 9i64);
- let source_pointer = unsafe { ObjectArc::as_raw(&source.data) };
- let calls_before = SHALLOW_COPY_CALLS.load(Ordering::Relaxed);
+fn reflected_no_change_returns_original() {
+ let source = reflected_object();
let mapped = structural_map(
source.clone(),
|string: FfiString| Any::from(string),
WalkOrder::PostOrder,
)
- .and_then(RustPair::try_from)
.unwrap();
- assert_eq!(unsafe { ObjectArc::as_raw(&mapped.data) }, source_pointer);
- assert_eq!(SHALLOW_COPY_CALLS.load(Ordering::Relaxed), calls_before + 1);
+ assert_eq!(any_object_pointer(&mapped), any_object_pointer(&source));
}
#[test]
fn reflected_object_without_shallow_copy_is_rejected_even_when_unchanged() {
- ensure_test_types_registered();
- let error = match structural_map(
- rust_no_copy(),
- |_integer: i64| Any::from(0i64),
- WalkOrder::PostOrder,
- ) {
- Ok(_) => panic!("reflected object without a shallow-copy hook
unexpectedly succeeded"),
- Err(error) => error,
- };
- assert!(error.message().contains("__ffi_shallow_copy__"));
-}
-
-#[test]
-fn reflected_getter_releases_partial_result_on_error() {
- let tracked = FfiString::from("a reference-counted reflected field value");
- let source = rust_failing_getter(tracked.clone());
- let count_before = AnyView::from(&tracked).debug_strong_count();
-
- let error = match structural_map(
- source.clone(),
- |_integer: i64| Any::from(0i64),
- WalkOrder::PostOrder,
- ) {
- Ok(_) => panic!("failing getter unexpectedly succeeded"),
- Err(error) => error,
- };
-
+ // Keep the C++ test library linked for its startup registrations.
assert_eq!(
- error.message(),
- "callback failed after writing an owning result"
+ unsafe { tvm_ffi::tvm_ffi_sys::TVMFFITestingDummyTarget() },
+ 0
);
- assert_eq!(AnyView::from(&tracked).debug_strong_count(), count_before);
-}
-
-#[test]
-fn function_setter_releases_partial_result_on_error() {
- let replacement = FfiString::from("a reference-counted setter result");
- let source = rust_failing_setter(1i64);
- let count_before = AnyView::from(&replacement).debug_strong_count();
-
+ // This existing C++ test type deletes its copy constructor.
+ let source = Function::from_type_key_method("testing.TestNonCopyable",
"__ffi_init__")
+ .unwrap()
+ .call_tuple((1i64,))
+ .unwrap();
+ // Leave its integer field unmatched to test the unchanged-object path.
let error = match structural_map(
source,
- |_value: i64| Any::from(replacement.clone()),
+ |string: FfiString| Any::from(string),
WalkOrder::PostOrder,
) {
- Ok(_) => panic!("failing Function setter unexpectedly succeeded"),
+ Ok(_) => panic!("reflected object without a shallow-copy hook
unexpectedly succeeded"),
Err(error) => error,
};
-
- assert_eq!(
- error.message(),
- "callback failed after writing an owning result"
- );
- assert_eq!(
- AnyView::from(&replacement).debug_strong_count(),
- count_before
- );
+ assert!(error.message().contains("__ffi_shallow_copy__"));
}
#[test]
fn callback_errors_preserve_message_and_add_object_context() {
- ensure_test_types_registered();
let error = match structural_map(
Array::new(vec![1i64]),
|_integer: i64| -> Result<i64> {
@@ -1103,105 +425,6 @@ fn
callback_errors_preserve_message_and_add_object_context() {
assert!(error.backtrace().contains("object `ffi.Array`"));
}
-#[test]
-fn registered_mutation_hooks_receive_the_rust_mutator() {
- ensure_test_types_registered();
- let _guard = REGISTERED_HOOK_TEST_LOCK.lock().unwrap();
- RETAINED_MUTATOR.with(|retained| {
- retained.take();
- });
-
- let source = rust_hook_node();
- let mutate_calls_before = REGISTERED_MUTATE_CALLS.load(Ordering::Relaxed);
- let mutated = structural_mutate(source.clone(), &mut
ManualIncrement::default())
- .and_then(i64::try_from)
- .unwrap();
- assert_eq!(mutated, 2);
- assert_eq!(
- REGISTERED_MUTATE_CALLS.load(Ordering::Relaxed),
- mutate_calls_before + 1
- );
-
- let retained = RETAINED_MUTATOR.with(|retained| retained.take().unwrap());
- let error = match Function::get_global("ffi.StructuralMutatorMutate")
- .unwrap()
- .call_packed(&[AnyView::from(&retained), AnyView::from(&1i64)])
- {
- Ok(_) => panic!("retained structural mutator unexpectedly remained
active"),
- Err(error) => error,
- };
- assert!(error.message().contains("retained after its active call"));
-
- let mutate_calls_before = REGISTERED_MUTATE_CALLS.load(Ordering::Relaxed);
- let mutated = structural_mutate(
- source.clone(),
- |value: i64, _mutator: &mut CallbackMutator| Any::from(value + 1),
- )
- .and_then(i64::try_from)
- .unwrap();
- assert_eq!(mutated, 2);
- assert_eq!(
- REGISTERED_MUTATE_CALLS.load(Ordering::Relaxed),
- mutate_calls_before + 1
- );
- RETAINED_MUTATOR.with(|retained| {
- retained.take();
- });
-
- let maybe_inplace_calls_before =
REGISTERED_MAYBE_INPLACE_MUTATE_CALLS.load(Ordering::Relaxed);
- let mutated = structural_mutate(rust_hook_node(), &mut
ManualIncrement::default())
- .and_then(i64::try_from)
- .unwrap();
- assert_eq!(mutated, 2);
- assert_eq!(
- REGISTERED_MAYBE_INPLACE_MUTATE_CALLS.load(Ordering::Relaxed),
- maybe_inplace_calls_before + 1
- );
- RETAINED_MUTATOR.with(|retained| {
- retained.take();
- });
-}
-
-#[test]
-fn registered_hook_cannot_reenter_through_an_aliased_mutator_handle() {
- ensure_test_types_registered();
- let _guard = REGISTERED_HOOK_TEST_LOCK.lock().unwrap();
- RETAINED_MUTATOR.with(|retained| {
- retained.take();
- });
- let mut mutator = RejectAliasedReentry::default();
-
- let mutated = structural_mutate(rust_hook_node(), &mut mutator)
- .and_then(i64::try_from)
- .unwrap();
-
- assert_eq!(mutated, 2);
- assert!(mutator.rejected);
- RETAINED_MUTATOR.with(|retained| {
- retained.take();
- });
-}
-
-#[test]
-fn registered_hook_rejects_foreign_thread_mutator_callback() {
- ensure_test_types_registered();
- let _guard = REGISTERED_HOOK_TEST_LOCK.lock().unwrap();
- FOREIGN_THREAD_MUTATOR_ERROR.with(|error| {
- error.take();
- });
-
- PROBE_FOREIGN_THREAD_MUTATOR.with(|enabled| enabled.set(true));
- let result = structural_mutate(rust_hook_node(), &mut
ManualIncrement::default());
- PROBE_FOREIGN_THREAD_MUTATOR.with(|enabled| enabled.set(false));
- result.unwrap();
-
- let message = FOREIGN_THREAD_MUTATOR_ERROR.with(|error|
error.take().unwrap());
- assert!(message.contains("invoked from a different thread"));
- RETAINED_MUTATOR.with(|retained| {
- retained.take();
- });
-}
-
#[test]
fn callback_panics_resume_after_the_registered_hook_returns() {
let panic = match std::panic::catch_unwind(std::panic::AssertUnwindSafe(||
{
@@ -1233,7 +456,6 @@ fn
callback_panics_resume_after_the_registered_hook_returns() {
#[test]
fn unique_map_reuses_nested_unique_value_storage() {
- ensure_test_types_registered();
let child = Array::new(vec![1i64, 2]);
let child_pointer = array_pointer(&child);
let source: Map<i64, Array<i64>> = [(1, child)].into_iter().collect();
@@ -1251,7 +473,6 @@ fn unique_map_reuses_nested_unique_value_storage() {
#[test]
fn shared_map_and_dict_copy_only_when_a_value_changes() {
- ensure_test_types_registered();
let source: Map<i64, i64> = [(1, 10)].into_iter().collect();
let source_pointer = map_pointer(&source);
let unchanged = structural_map(
@@ -1289,7 +510,6 @@ fn shared_map_and_dict_copy_only_when_a_value_changes() {
#[test]
fn shared_map_callback_error_preserves_source_and_reports_object_context() {
- ensure_test_types_registered();
let source: Map<i64, i64> = [(1, 10), (2, 20)].into_iter().collect();
let error = match structural_map(
source.clone(),
@@ -1311,7 +531,6 @@ fn
shared_map_callback_error_preserves_source_and_reports_object_context() {
#[test]
fn shared_outer_container_does_not_mutate_its_nested_child() {
- ensure_test_types_registered();
let nested = call_global("ffi.List", &[Any::from(1i64)]);
let nested_pointer = any_object_pointer(&nested);
let outer = call_global("ffi.Array", &[nested]);
@@ -1330,7 +549,6 @@ fn
shared_outer_container_does_not_mutate_its_nested_child() {
#[test]
fn shared_list_uses_copy_on_write() {
- ensure_test_types_registered();
let source = call_global("ffi.List", &[Any::from(1i64), Any::from(2i64)]);
let source_pointer = any_object_pointer(&source);
let mapped =
@@ -1362,7 +580,6 @@ impl GeneratedMapper {
#[test]
fn generated_map_dispatch_supports_kind_and_ordered_catch_all() {
- ensure_test_types_registered();
let mut mapper = GeneratedMapper::default();
let mapped = structural_map(Array::new(vec![1i64, 2]), &mut mapper,
WalkOrder::PostOrder)
.and_then(Array::<i64>::try_from)
@@ -1426,25 +643,6 @@ fn
generated_mutate_dispatch_defaults_unmatched_values_and_preserves_inplace_per
);
}
-#[test]
-fn generated_mutate_dispatch_default_remap_crosses_registered_hooks() {
- ensure_test_types_registered();
- let _guard = REGISTERED_HOOK_TEST_LOCK.lock().unwrap();
- RETAINED_MUTATOR.with(|retained| {
- retained.take();
- });
-
- let mut mutator = GeneratedLeafDispatch::default();
- let mutated = structural_mutate(rust_hook_node(), &mut mutator)
- .and_then(i64::try_from)
- .unwrap();
- assert_eq!(mutated, 2);
- assert_eq!(mutator.integers, vec![(1, DefRegionKind::None)]);
- RETAINED_MUTATOR.with(|retained| {
- retained.take();
- });
-}
-
#[derive(Default)]
struct GeneratedRecursiveDispatch {
arrays: Vec<DefRegionKind>,
@@ -1485,23 +683,6 @@ fn generated_mutate_dispatch_recurses_through_context() {
);
}
-#[test]
-fn generated_mutate_dispatch_inherits_region_during_explicit_recursion() {
- ensure_test_types_registered();
- let _guard = REFLECTED_TEST_LOCK.lock().unwrap();
- let root = rust_pair(Array::new(vec![1i64]), Any::new());
- let mut mutator = GeneratedRecursiveDispatch::default();
-
- let mutated = structural_mutate(root, &mut mutator)
- .and_then(RustPair::try_from)
- .unwrap();
- let first = Array::<i64>::try_from(mutated.data.first.clone()).unwrap();
-
- assert_eq!(first.iter().collect::<Vec<_>>(), vec![11]);
- assert_eq!(mutator.arrays, vec![DefRegionKind::Recursive]);
- assert_eq!(mutator.integers, vec![(1, DefRegionKind::Recursive)]);
-}
-
#[derive(Default)]
struct GeneratedDefaultingDispatch {
arrays: usize,
@@ -1533,55 +714,8 @@ fn
generated_mutate_dispatch_can_default_recurse_from_a_typed_handler() {
assert_eq!(mutator.integers, vec![1, 2]);
}
-struct GeneratedRemappingDispatch {
- type_index: i32,
- calls: usize,
-}
-
-#[dispatch(mutate)]
-impl GeneratedRemappingDispatch {
- fn mutate_dag_node(&mut self, _value: &RustDagNodeObj) -> Any {
- Any::from(42i64)
- }
-
- fn mutate_any(&mut self, value: &MapValue, mutator: &mut Mutator) ->
Result<Any> {
- if value.type_index() != self.type_index {
- return mutator.default_mutate(self);
- }
- if let Some(mutated) = mutator.var_remap_get(self, value)? {
- return Ok(mutated);
- }
- self.calls += 1;
- let mutated = Any::from(41i64);
- mutator.var_remap_set(self, value, &mutated)?;
- Ok(mutated)
- }
-}
-
-#[test]
-fn generated_mutate_dispatch_uses_fresh_invocation_local_var_remap() {
- ensure_test_types_registered();
- let var = rust_free_var();
- let mut mutator = GeneratedRemappingDispatch {
- type_index: RustFreeVarObj::type_index(),
- calls: 0,
- };
-
- for expected_calls in [1, 2] {
- let root = call_global(
- "ffi.Array",
- &[Any::from(var.clone()), Any::from(var.clone())],
- );
- let mutated = structural_mutate(root, &mut mutator).unwrap();
- assert_eq!(mutator.calls, expected_calls);
- assert_eq!(i64::try_from(array_item(&mutated, 0)).unwrap(), 41);
- assert_eq!(i64::try_from(array_item(&mutated, 1)).unwrap(), 41);
- }
-}
-
#[test]
fn pre_order_retained_alias_disables_in_place_mutation() {
- ensure_test_types_registered();
let root = call_global("ffi.List", &[Any::from(1i64)]);
let root_pointer = any_object_pointer(&root);
let mut retained = None;
@@ -1610,7 +744,6 @@ fn pre_order_retained_alias_disables_in_place_mutation() {
#[test]
fn closures_and_tuples_use_ordered_first_match() {
- ensure_test_types_registered();
let root = Array::new(vec![1i64, 2]);
let mapped = structural_map(
root,
@@ -1646,7 +779,6 @@ fn closures_and_tuples_use_ordered_first_match() {
#[test]
fn callbacks_return_values_convertible_into_any() {
- ensure_test_types_registered();
let mapped = structural_map(
Array::new(vec![1i64, 2]),
|integer: i64| integer + 10,
@@ -1684,8 +816,8 @@ fn twelve_link_tuple_reaches_final_map_dispatch() {
|_value: f64| Any::from(0.0f64),
|value: FfiString| Any::from(value),
|value: Function| Any::from(value),
- |_node: &RustDagNodeObj| Any::new(),
- |_node: &RustFreeVarObj| Any::new(),
+ |_node: &MapObj| Any::new(),
+ |_node: &FunctionObj| Any::new(),
|value: Array<i64>| Any::from(value),
|value: Array<f64>| Any::from(value),
|value: Array<bool>| Any::from(value),
@@ -1703,7 +835,6 @@ fn twelve_link_tuple_reaches_final_map_dispatch() {
#[test]
fn nested_tuple_chain_exceeds_flat_arity() {
- ensure_test_types_registered();
let root = Array::new(vec![1i64, 2, 3]);
let mut catch_all = 0;
let mapped = structural_map(
@@ -1714,8 +845,8 @@ fn nested_tuple_chain_exceeds_flat_arity() {
|_value: f64| Any::from(0.0f64),
|value: FfiString| Any::from(value),
|value: Function| Any::from(value),
- |_node: &RustDagNodeObj| Any::new(),
- |_node: &RustFreeVarObj| Any::new(),
+ |_node: &MapObj| Any::new(),
+ |_node: &FunctionObj| Any::new(),
|value: Array<f64>| Any::from(value),
|value: Array<bool>| Any::from(value),
|value: Array<Array<i64>>| Any::from(value),
@@ -1744,7 +875,6 @@ fn nested_tuple_chain_exceeds_flat_arity() {
#[test]
fn callbacks_run_in_the_configured_order() {
- ensure_test_types_registered();
let root = Array::new(vec![1i64, 2]);
let mut pre = Vec::new();
structural_map(
@@ -1773,7 +903,6 @@ fn callbacks_run_in_the_configured_order() {
#[test]
fn map_keys_are_anchors_and_object_leaves_are_preserved() {
- ensure_test_types_registered();
let root: Map<FfiString, i64> = [(FfiString::from("a"), 1i64),
(FfiString::from("b"), 2i64)]
.into_iter()
.collect();
@@ -1823,7 +952,6 @@ fn map_keys_are_anchors_and_object_leaves_are_preserved() {
#[test]
fn callback_mutate_defaults_unmatched_values_and_preserves_root_permit() {
- ensure_test_types_registered();
let root = Array::new(vec![1i64, 2]);
let root_pointer = array_pointer(&root);
let mutated = structural_mutate(root, |value: i64, _mutator: &mut
CallbackMutator| {
@@ -1856,7 +984,6 @@ fn stateful_mutate_default(
#[test]
fn callback_mutator_carries_reusable_mutable_state() {
- ensure_test_types_registered();
let mut mutator = MutateCallbacks::new(
CallbackMutateStats::default(),
(stateful_mutate_integer, stateful_mutate_default),
@@ -1907,7 +1034,6 @@ fn stateful_mutate_recursive(
#[test]
fn callback_mutator_state_can_change_around_recursive_reborrows() {
- ensure_test_types_registered();
let root = Array::new(vec![Array::new(vec![1i64, 2])]);
let mut mutator =
MutateCallbacks::new(CallbackMutateDepth::default(),
stateful_mutate_recursive);
@@ -1926,7 +1052,6 @@ fn
callback_mutator_state_can_change_around_recursive_reborrows() {
#[test]
fn callback_mutate_current_default_is_repeatable_copy_path() {
- ensure_test_types_registered();
let root = Array::new(vec![1i64, 2]);
let root_pointer = array_pointer(&root);
let defaults = Cell::new(0);
@@ -1952,7 +1077,6 @@ fn
callback_mutate_current_default_is_repeatable_copy_path() {
#[test]
fn callback_mutate_match_is_final_and_same_fn_can_reenter() {
- ensure_test_types_registered();
let integer_calls = Cell::new(0);
let mutated = structural_mutate(
Array::new(vec![1i64]),
@@ -1985,19 +1109,20 @@ fn
callback_mutate_match_is_final_and_same_fn_can_reenter() {
#[test]
fn callback_mutate_supports_node_links_nested_tuples_and_reflection() {
- ensure_test_types_registered();
- let _guard = REFLECTED_TEST_LOCK.lock().unwrap();
let root = call_global(
"ffi.Array",
- &[Any::from(rust_dag_node()), Any::from(rust_pair(1i64, 9i64))],
+ &[
+ Any::from(Function::from_packed(|_| Ok(Any::new()))),
+ reflected_object(),
+ ],
);
let regions = RefCell::new(Vec::new());
let mutated = structural_mutate(
root,
(
(
- |_value: f64, _mutator: &mut CallbackMutator| Any::new(),
- |_node: &RustDagNodeObj, _mutator: &mut CallbackMutator|
Any::from(7i64),
+ |_value: bool, _mutator: &mut CallbackMutator| Any::new(),
+ |_node: &FunctionObj, _mutator: &mut CallbackMutator|
Any::from(7i64),
),
|value: i64, mutator: &mut CallbackMutator| {
regions.borrow_mut().push(mutator.def_region_kind());
@@ -2008,15 +1133,18 @@ fn
callback_mutate_supports_node_links_nested_tuples_and_reflection() {
.unwrap();
assert_eq!(i64::try_from(array_item(&mutated, 0)).unwrap(), 7);
- let pair = RustPair::try_from(array_item(&mutated, 1)).unwrap();
- assert_eq!(i64::try_from(pair.data.first.clone()).unwrap(), 2);
- assert_eq!(i64::try_from(pair.data.ignored.clone()).unwrap(), 9);
- assert_eq!(*regions.borrow(), vec![DefRegionKind::Recursive]);
+ let object = array_item(&mutated, 1);
+ assert_eq!(reflected_field::<i64>(&object, "v_i64"), 2);
+ assert_eq!(reflected_field::<f64>(&object, "v_f64"), 2.5);
+ assert_eq!(
+ reflected_field::<FfiString>(&object, "v_str").as_str(),
+ "a reflected string"
+ );
+ assert_eq!(*regions.borrow(), vec![DefRegionKind::None]);
}
#[test]
fn callback_mutate_distinguishes_borrowed_and_owned_children() {
- ensure_test_types_registered();
let borrowed_child = Array::new(vec![1i64]);
let borrowed_pointer = array_pointer(&borrowed_child);
let mutated = structural_mutate(
@@ -2050,40 +1178,6 @@ fn
callback_mutate_distinguishes_borrowed_and_owned_children() {
assert_eq!(mutated.get(0).unwrap(), 2);
}
-#[test]
-fn callback_mutate_can_use_its_invocation_local_var_remap() {
- ensure_test_types_registered();
- let var = rust_free_var();
- let calls = Cell::new(0);
- let type_index = RustFreeVarObj::type_index();
- let mut mutator = MutateCallbacks::new(
- (),
- |value: &MapValue, mutator: &mut CallbackMutator| -> Result<Any> {
- if value.type_index() != type_index {
- return mutator.default_mutate();
- }
- if let Some(mutated) = mutator.var_remap_get(value)? {
- return Ok(mutated);
- }
- calls.set(calls.get() + 1);
- let mutated = Any::from(41i64);
- mutator.var_remap_set(value, &mutated)?;
- Ok(mutated)
- },
- );
-
- for expected_calls in 1..=2 {
- let root = call_global(
- "ffi.Array",
- &[Any::from(var.clone()), Any::from(var.clone())],
- );
- let mutated = structural_mutate(root, &mut mutator).unwrap();
- assert_eq!(calls.get(), expected_calls);
- assert_eq!(i64::try_from(array_item(&mutated, 0)).unwrap(), 41);
- assert_eq!(i64::try_from(array_item(&mutated, 1)).unwrap(), 41);
- }
-}
-
#[test]
fn nested_callback_mutate_restores_the_outer_active_mutator() {
let mutated = structural_mutate(
diff --git a/rust/tvm-ffi/tests/test_structural_visit.rs
b/rust/tvm-ffi/tests/test_structural_visit.rs
index a11267f4..8a141cd7 100644
--- a/rust/tvm-ffi/tests/test_structural_visit.rs
+++ b/rust/tvm-ffi/tests/test_structural_visit.rs
@@ -18,353 +18,56 @@
*/
use std::cell::{Cell, RefCell};
-use std::sync::LazyLock;
-
-use tvm_ffi::derive::{Object as DeriveObject, ObjectRef as DeriveObjectRef};
use tvm_ffi::object::ObjectRef;
-use tvm_ffi::tvm_ffi_sys::{
- TVMFFIAny, TVMFFIAnyViewToOwnedAny, TVMFFIByteArray,
TVMFFIFieldFlagBitMask, TVMFFIFieldInfo,
- TVMFFISEqHashKind, TVMFFITypeMetadata, TVMFFITypeRegisterAttr,
-};
use tvm_ffi::{
- dispatch, get_type_attr, structural_visit, structural_walk, Any, AnyView,
Array, DLDataType,
- DLDataTypeCode, DefRegionKind, Error, FieldGetter, Function, Map, Object,
ObjectArc,
- ObjectCore, ObjectRefCast, Result, String as FfiString, StructuralVisitor,
TypeIndex,
- VisitCallbacks, VisitContext, VisitInterrupt, VisitValue, WalkOrder,
WalkResult, RUNTIME_ERROR,
+ dispatch, get_type_attr, structural_visit, structural_walk, Any, Array,
DLDataType,
+ DLDataTypeCode, DefRegionKind, Error, FieldGetter, Function, Map, Object,
ObjectRefCore,
+ Result, String as FfiString, StructuralVisitor, TypeIndex, VisitCallbacks,
VisitContext,
+ VisitInterrupt, VisitValue, WalkOrder, WalkResult, RUNTIME_ERROR,
};
-unsafe extern "C" {
- fn TVMFFITypeGetOrAllocIndex(
- type_key: *const TVMFFIByteArray,
- static_type_index: i32,
- type_depth: i32,
- num_child_slots: i32,
- child_slots_can_overflow: i32,
- parent_type_index: i32,
- ) -> i32;
- fn TVMFFITypeRegisterField(type_index: i32, info: *const TVMFFIFieldInfo)
-> i32;
- fn TVMFFITypeRegisterMetadata(type_index: i32, metadata: *const
TVMFFITypeMetadata) -> i32;
-}
-
-#[repr(C)]
-#[derive(DeriveObject)]
-#[type_key = "testing.RustStructuralVisitDefRegion"]
-#[type_final]
-struct RustVisitDefRegionObj {
- base: Object,
- recursive: Any,
- plain: Any,
- non_recursive: Any,
- both: Any,
- ignored: Any,
-}
-
-#[repr(C)]
-#[derive(DeriveObjectRef, Clone)]
-struct RustVisitDefRegion {
- data: ObjectArc<RustVisitDefRegionObj>,
-}
-
-#[repr(C)]
-#[derive(DeriveObject)]
-#[type_key = "testing.RustStructuralVisitHook"]
-#[type_final]
-struct RustVisitHookObj {
- base: Object,
- selected: Any,
- ignored: Any,
-}
-
-#[repr(C)]
-#[derive(DeriveObjectRef, Clone)]
-struct RustVisitHook {
- data: ObjectArc<RustVisitHookObj>,
-}
-
-#[repr(C)]
-#[derive(DeriveObject)]
-#[type_key = "testing.RustStructuralVisitFailingGetter"]
-#[type_final]
-struct RustVisitFailingGetterObj {
- base: Object,
- value: Any,
-}
-
-#[repr(C)]
-#[derive(DeriveObjectRef, Clone)]
-struct RustVisitFailingGetter {
- data: ObjectArc<RustVisitFailingGetterObj>,
-}
-
-thread_local! {
- static RETAINED_VISITOR: RefCell<Option<Any>> = const { RefCell::new(None)
};
- static REGISTERED_HOOK_REGIONS: RefCell<Vec<i64>> = const {
RefCell::new(Vec::new()) };
- static PROBE_FOREIGN_THREAD_VISITOR: Cell<bool> = const { Cell::new(false)
};
- static FOREIGN_THREAD_VISITOR_ERROR: RefCell<Option<String>> = const {
RefCell::new(None) };
-}
-
-unsafe extern "C" fn clone_any_field(field: *mut std::ffi::c_void, result:
*mut TVMFFIAny) -> i32 {
- TVMFFIAnyViewToOwnedAny(field.cast(), result)
-}
-
-unsafe extern "C" fn clone_any_field_then_fail(
- field: *mut std::ffi::c_void,
- result: *mut TVMFFIAny,
-) -> i32 {
- let code = TVMFFIAnyViewToOwnedAny(field.cast(), result);
- if code != 0 {
- return code;
- }
- Error::set_raised(&runtime_error(
- "visit getter failed after writing an owning result",
- ));
- -1
-}
-
-fn register_any_field(type_index: i32, name: &'static str, offset: usize,
flags: i64) {
- let field = TVMFFIFieldInfo {
- name: unsafe { TVMFFIByteArray::from_str(name) },
- doc: unsafe { TVMFFIByteArray::from_str("Rust structural-visit test
field") },
- metadata: unsafe { TVMFFIByteArray::from_str("") },
- flags,
- size: std::mem::size_of::<Any>() as i64,
- alignment: std::mem::align_of::<Any>() as i64,
- offset: offset as i64,
- getter: Some(clone_any_field),
- setter: std::ptr::null_mut(),
- default_value_or_factory: TVMFFIAny::new(),
- field_static_type_index: -1,
- };
- assert_eq!(unsafe { TVMFFITypeRegisterField(type_index, &field) }, 0);
-}
-
-fn register_visit_type(type_key: &'static str, total_size: usize, kind:
TVMFFISEqHashKind) -> i32 {
- let type_key = unsafe { TVMFFIByteArray::from_str(type_key) };
- let type_index = unsafe {
- TVMFFITypeGetOrAllocIndex(
- &type_key,
- -1,
- Object::TYPE_DEPTH + 1,
- 0,
- 1,
- Object::type_index(),
- )
- };
- assert!(type_index >= TypeIndex::kTVMFFIDynObjectBegin as i32);
- let metadata = TVMFFITypeMetadata {
- doc: unsafe { TVMFFIByteArray::from_str("Rust structural-visit test
object") },
- creator: None,
- total_size: i32::try_from(total_size).unwrap(),
- structural_eq_hash_kind: kind as i32,
- };
- assert_eq!(
- unsafe { TVMFFITypeRegisterMetadata(type_index, &metadata) },
- 0
- );
- type_index
-}
-
-fn register_function_attr(type_index: i32, name: &'static str, function:
Function) {
- let name = unsafe { TVMFFIByteArray::from_str(name) };
- let mut value = Any::from(function);
- assert_eq!(
- unsafe { TVMFFITypeRegisterAttr(type_index, &name,
Any::as_data_ptr(&mut value)) },
- 0
- );
-}
-
-fn call_visitor_from_foreign_thread(visitor: AnyView<'_>) -> String {
- // Keep the object alive on this thread. The worker uses a non-owning view,
- // so it neither transfers nor releases the visitor's reference count.
- let mut owner = Any::from(visitor);
- let raw = unsafe { *Any::as_data_ptr(&mut owner) };
- let type_index = raw.type_index;
- let object = unsafe { raw.data_union.v_obj } as usize;
- std::thread::spawn(move || {
- let mut raw = TVMFFIAny::new();
- raw.type_index = type_index;
- raw.data_union.v_obj = object as *mut _;
- let borrowed = std::mem::ManuallyDrop::new(unsafe {
Any::from_raw_ffi_any(raw) });
- match Function::get_global("ffi.StructuralVisitorVisit")
- .unwrap()
- .call_packed(&[AnyView::from(&*borrowed), AnyView::from(&1i64)])
- {
- Err(error) => error.message().to_string(),
- Ok(_) => "foreign-thread visitor call unexpectedly
succeeded".to_string(),
- }
- })
- .join()
- .unwrap()
-}
-
-fn registered_visit_hook(args: &[AnyView<'_>]) -> Result<Any> {
- assert_eq!(args.len(), 2);
- if PROBE_FOREIGN_THREAD_VISITOR.with(Cell::get) {
- let message = call_visitor_from_foreign_thread(args[0]);
- FOREIGN_THREAD_VISITOR_ERROR.with(|error|
error.replace(Some(message)));
- }
- RETAINED_VISITOR.with(|retained| {
- retained.replace(Some(Any::from(args[0])));
- });
- let def_region_kind =
Function::get_global("ffi.StructuralVisitorDefRegionKind")?
- .call_packed(&[args[0]])
- .and_then(i64::try_from)?;
- REGISTERED_HOOK_REGIONS.with(|regions|
regions.borrow_mut().push(def_region_kind));
-
- let node = RustVisitHook::try_from(args[1])?;
- Function::get_global("ffi.StructuralVisitorVisit")?
- .call_packed(&[args[0], AnyView::from(&node.data.selected)])
-}
-
-fn registered_primitive_visit_hook(args: &[AnyView<'_>]) -> Result<Any> {
- assert_eq!(args.len(), 2);
- Function::get_global("ffi.StructuralVisitorVisit")?
- .call_packed(&[args[0], AnyView::from(&7i64)])
-}
-
-// The runtime type table leaves registration synchronization to its callers.
-// Every test, including those using only built-in types, must wait for this
-// initializer before reading reflection data: registering a custom __s_visit__
-// hook can relocate the column containing the built-in Array/Map hooks.
-static REGISTER_TEST_TYPES: LazyLock<()> = LazyLock::new(|| {
- let type_index = register_visit_type(
- RustVisitHookObj::TYPE_KEY,
- std::mem::size_of::<RustVisitHookObj>(),
- TVMFFISEqHashKind::kTVMFFISEqHashKindTreeNode,
- );
- register_any_field(
- type_index,
- "selected",
- std::mem::offset_of!(RustVisitHookObj, selected),
- 0,
- );
- register_any_field(
- type_index,
- "ignored",
- std::mem::offset_of!(RustVisitHookObj, ignored),
- 0,
- );
- register_function_attr(
- type_index,
- "__s_visit__",
- Function::from_packed(registered_visit_hook),
- );
-
- let type_index = register_visit_type(
- RustVisitFailingGetterObj::TYPE_KEY,
- std::mem::size_of::<RustVisitFailingGetterObj>(),
- TVMFFISEqHashKind::kTVMFFISEqHashKindTreeNode,
- );
- let field = TVMFFIFieldInfo {
- name: unsafe { TVMFFIByteArray::from_str("value") },
- doc: unsafe { TVMFFIByteArray::from_str("Fail after producing an
owning field value") },
- metadata: unsafe { TVMFFIByteArray::from_str("") },
- flags: 0,
- size: std::mem::size_of::<Any>() as i64,
- alignment: std::mem::align_of::<Any>() as i64,
- offset: std::mem::offset_of!(RustVisitFailingGetterObj, value) as i64,
- getter: Some(clone_any_field_then_fail),
- setter: std::ptr::null_mut(),
- default_value_or_factory: TVMFFIAny::new(),
- field_static_type_index: -1,
- };
- assert_eq!(unsafe { TVMFFITypeRegisterField(type_index, &field) }, 0);
-
- register_function_attr(
- TypeIndex::kTVMFFIDataType as i32,
- "__s_visit__",
- Function::from_packed(registered_primitive_visit_hook),
- );
-
- let type_index = register_visit_type(
- RustVisitDefRegionObj::TYPE_KEY,
- std::mem::size_of::<RustVisitDefRegionObj>(),
- TVMFFISEqHashKind::kTVMFFISEqHashKindFreeVar,
- );
- for (name, offset, flags) in [
- (
- "recursive",
- std::mem::offset_of!(RustVisitDefRegionObj, recursive),
- TVMFFIFieldFlagBitMask::kTVMFFIFieldFlagBitMaskSEqHashDefRecursive
as i64,
- ),
- (
- "plain",
- std::mem::offset_of!(RustVisitDefRegionObj, plain),
- 0,
- ),
- (
- "non_recursive",
- std::mem::offset_of!(RustVisitDefRegionObj, non_recursive),
-
TVMFFIFieldFlagBitMask::kTVMFFIFieldFlagBitMaskSEqHashDefNonRecursive as i64,
- ),
- (
- "both",
- std::mem::offset_of!(RustVisitDefRegionObj, both),
- TVMFFIFieldFlagBitMask::kTVMFFIFieldFlagBitMaskSEqHashDefRecursive
as i64
- |
TVMFFIFieldFlagBitMask::kTVMFFIFieldFlagBitMaskSEqHashDefNonRecursive as i64,
- ),
- (
- "ignored",
- std::mem::offset_of!(RustVisitDefRegionObj, ignored),
- TVMFFIFieldFlagBitMask::kTVMFFIFieldFlagBitMaskSEqHashIgnore as
i64,
- ),
- ] {
- register_any_field(type_index, name, offset, flags);
- }
-});
-
-fn test_prelude() {
- LazyLock::force(®ISTER_TEST_TYPES);
-}
-
-fn rust_visit_hook(selected: impl Into<Any>, ignored: impl Into<Any>) ->
RustVisitHook {
- test_prelude();
- RustVisitHook {
- data: ObjectArc::new(RustVisitHookObj {
- base: Object::new(),
- selected: selected.into(),
- ignored: ignored.into(),
- }),
- }
-}
-
-fn rust_visit_failing_getter(value: impl Into<Any>) -> RustVisitFailingGetter {
- test_prelude();
- RustVisitFailingGetter {
- data: ObjectArc::new(RustVisitFailingGetterObj {
- base: Object::new(),
- value: value.into(),
- }),
- }
-}
-
fn runtime_error(message: &str) -> Error {
Error::new(RUNTIME_ERROR, message, "")
}
#[test]
fn public_reflection_access_uses_registered_field_and_type_attr() {
- test_prelude();
- let root = rust_visit_hook(FfiString::from("owned field"), 99i64);
- let type_index = RustVisitHookObj::type_index();
-
- let getter = FieldGetter::new(type_index, "selected").unwrap();
- let selected = getter.get::<_, FfiString>(&*root.data).unwrap();
+ // Keep the existing C++ test library linked so its startup registrations
+ // are available even when this test is run by itself.
+ assert_eq!(
+ unsafe { tvm_ffi::tvm_ffi_sys::TVMFFITestingDummyTarget() },
+ 0
+ );
+ let root = Function::get_global("ffi.MakeObjectFromPackedArgs")
+ .unwrap()
+ .call_tuple((
+ FfiString::from("testing.TestObjectBase"),
+ FfiString::from("v_str"),
+ FfiString::from("an owning C++ reflected field value"),
+ ))
+ .unwrap();
+ let type_index = root.type_index();
+ let root = ObjectRef::try_from(root).unwrap();
+
+ let getter = FieldGetter::new(type_index, "v_str").unwrap();
+ let selected = getter
+ .get::<_, FfiString>(&**ObjectRef::data(&root))
+ .unwrap();
drop(root);
- assert_eq!(selected.as_str(), "owned field");
+ assert_eq!(selected.as_str(), "an owning C++ reflected field value");
- let wrong_type = rust_visit_failing_getter(0i64);
- assert!(getter.get_any(&*wrong_type.data).is_err());
+ let wrong_type = Array::new(vec![0i64]);
+ assert!(getter.get_any(&**Array::data(&wrong_type)).is_err());
assert!(FieldGetter::new(type_index, "missing").is_err());
- assert!(Function::try_from(get_type_attr(type_index,
"__s_visit__").unwrap()).is_ok());
- assert!(Function::from_type_attr(type_index, "__s_visit__").is_ok());
+ // ObjectDef registers this Function-valued attribute for copyable C++
types.
+ assert!(Function::try_from(get_type_attr(type_index,
"__ffi_shallow_copy__").unwrap()).is_ok());
+ assert!(Function::from_type_attr(type_index,
"__ffi_shallow_copy__").is_ok());
assert!(get_type_attr(type_index, "missing").is_none());
}
#[test]
fn plain_walk_uses_registered_array_hook() {
- test_prelude();
let root = Array::new(vec![1i64, 2, 3]);
let mut integers = 0;
assert!(structural_walk(
@@ -382,32 +85,8 @@ fn plain_walk_uses_registered_array_hook() {
assert_eq!(integers, 3);
}
-#[test]
-fn reflected_getter_releases_partial_result_on_error() {
- test_prelude();
- let tracked = FfiString::from("a reference-counted reflected visit field");
- let root = rust_visit_failing_getter(tracked.clone());
- let count_before = AnyView::from(&tracked).debug_strong_count();
-
- let error = match structural_walk(
- &root,
- |_value: &VisitValue| WalkResult::Advance,
- WalkOrder::PreOrder,
- ) {
- Ok(_) => panic!("failing getter unexpectedly succeeded"),
- Err(error) => error,
- };
-
- assert_eq!(
- error.message(),
- "visit getter failed after writing an owning result"
- );
- assert_eq!(AnyView::from(&tracked).debug_strong_count(), count_before);
-}
-
#[test]
fn plain_walk_visits_map_values_without_visiting_keys() {
- test_prelude();
let root: Map<FfiString, i64> = [(FfiString::from("a"), 1i64),
(FfiString::from("b"), 2i64)]
.into_iter()
.collect();
@@ -432,142 +111,8 @@ fn plain_walk_visits_map_values_without_visiting_keys() {
}
#[test]
-fn registered_function_hook_controls_children_interrupts_and_lifetime() {
- test_prelude();
- RETAINED_VISITOR.with(|retained| {
- retained.take();
- });
- REGISTERED_HOOK_REGIONS.with(|regions| regions.borrow_mut().clear());
- let root = rust_visit_hook(11i64, 99i64);
-
- let mut integers = Vec::new();
- assert!(structural_walk(
- &root,
- |value: &VisitValue| {
- if let Some(integer) = value.cast::<i64>() {
- integers.push(integer);
- }
- WalkResult::Advance
- },
- WalkOrder::PreOrder,
- )
- .unwrap()
- .is_none());
- // Reflection would visit both fields. The registered hook deliberately
- // visits only `selected`.
- assert_eq!(integers, vec![11]);
- REGISTERED_HOOK_REGIONS.with(|regions| {
- assert_eq!(regions.borrow().as_slice(), &[DefRegionKind::None as i64]);
- });
-
- #[derive(Default)]
- struct RecordingVisitor {
- integers: Vec<i64>,
- }
- impl StructuralVisitor for RecordingVisitor {
- fn visit(
- &mut self,
- value: &VisitValue,
- def_region_kind: DefRegionKind,
- ) -> Result<Option<VisitInterrupt>> {
- if let Some(integer) = value.cast::<i64>() {
- self.integers.push(integer);
- }
- self.default_visit_children(value, def_region_kind)
- }
- }
- let mut visitor = RecordingVisitor::default();
- assert!(structural_visit(&root, &mut visitor).unwrap().is_none());
- assert_eq!(visitor.integers, vec![11]);
-
- let callback_integers = RefCell::new(Vec::new());
- assert!(
- structural_visit(&root, |value: i64, _visitor: &mut VisitContext<'_,
()>| {
- callback_integers.borrow_mut().push(value)
- },)
- .unwrap()
- .is_none()
- );
- assert_eq!(*callback_integers.borrow(), vec![11]);
-
- test_prelude();
- let wrapped = RustVisitDefRegion {
- data: ObjectArc::new(RustVisitDefRegionObj {
- base: Object::new(),
- recursive: Any::from(root.clone()),
- plain: Any::new(),
- non_recursive: Any::new(),
- both: Any::new(),
- ignored: Any::new(),
- }),
- };
- assert!(structural_walk(
- &wrapped,
- |_value: &VisitValue| WalkResult::Advance,
- WalkOrder::PreOrder,
- )
- .unwrap()
- .is_none());
- REGISTERED_HOOK_REGIONS.with(|regions| {
- assert_eq!(
- regions.borrow().last().copied(),
- Some(DefRegionKind::Recursive as i64)
- );
- });
-
- let outcome = structural_walk(
- &root,
- |value: &VisitValue| match value.cast::<i64>() {
- Some(11) => WalkResult::interrupt_with(FfiString::from("stop")),
- _ => WalkResult::Advance,
- },
- WalkOrder::PreOrder,
- )
- .unwrap()
- .unwrap();
- assert_eq!(FfiString::try_from(outcome.value).unwrap().as_str(), "stop");
-
- let retained = RETAINED_VISITOR.with(|retained| retained.take().unwrap());
- let error = match Function::get_global("ffi.StructuralVisitorVisit")
- .unwrap()
- .call_packed(&[AnyView::from(&retained), AnyView::from(&1i64)])
- {
- Err(error) => error,
- Ok(_) => panic!("retained structural visitor unexpectedly remained
active"),
- };
- assert!(error.message().contains("retained after its active call"));
-}
-
-#[test]
-fn registered_hook_rejects_foreign_thread_visitor_callback() {
- test_prelude();
- RETAINED_VISITOR.with(|retained| {
- retained.take();
- });
- FOREIGN_THREAD_VISITOR_ERROR.with(|error| {
- error.take();
- });
-
- let root = rust_visit_hook(11i64, 99i64);
- PROBE_FOREIGN_THREAD_VISITOR.with(|enabled| enabled.set(true));
- let result = structural_walk(
- &root,
- |_value: &VisitValue| WalkResult::Advance,
- WalkOrder::PreOrder,
- );
- PROBE_FOREIGN_THREAD_VISITOR.with(|enabled| enabled.set(false));
- assert!(result.unwrap().is_none());
-
- let message = FOREIGN_THREAD_VISITOR_ERROR.with(|error|
error.take().unwrap());
- assert!(message.contains("invoked from a different thread"));
- RETAINED_VISITOR.with(|retained| {
- retained.take();
- });
-}
-
-#[test]
-fn primitive_hook_fast_path_preserves_pre_and_post_order() {
- test_prelude();
+fn primitive_values_are_leaves_in_pre_and_post_order() {
+ assert!(get_type_attr(TypeIndex::kTVMFFIDataType as i32,
"__s_visit__").is_none());
let dtype = DLDataType::new(DLDataTypeCode::kDLFloat, 32, 1);
let mut pre = Vec::new();
@@ -585,7 +130,7 @@ fn primitive_hook_fast_path_preserves_pre_and_post_order() {
)
.unwrap()
.is_none());
- assert_eq!(pre, ["dtype", "child"]);
+ assert_eq!(pre, ["dtype"]);
let mut skipped = Vec::new();
assert!(structural_walk(
@@ -620,12 +165,11 @@ fn
primitive_hook_fast_path_preserves_pre_and_post_order() {
)
.unwrap()
.is_none());
- assert_eq!(post, ["child", "dtype"]);
+ assert_eq!(post, ["dtype"]);
}
#[test]
fn primitive_fast_path_preserves_none_interrupt_and_error() {
- test_prelude();
let mut none_calls = 0;
assert!(structural_walk(
&Any::new(),
@@ -661,7 +205,6 @@ fn primitive_fast_path_preserves_none_interrupt_and_error()
{
#[test]
fn registered_map_hook_visits_all_values_without_visiting_keys() {
- test_prelude();
// More than 4 entries forces the dense (block + iteration list) layout.
let root: Map<FfiString, i64> = (0..9)
.map(|i| (FfiString::from(format!("k{i}")), i as i64))
@@ -688,7 +231,6 @@ fn
registered_map_hook_visits_all_values_without_visiting_keys() {
#[test]
fn interrupt_payload_crosses_map_traversal() {
- test_prelude();
let root: Map<FfiString, i64> = [(FfiString::from("a"), 1i64),
(FfiString::from("b"), 2i64)]
.into_iter()
.collect();
@@ -711,7 +253,6 @@ fn interrupt_payload_crosses_map_traversal() {
#[test]
fn handler_error_crosses_map_traversal() {
- test_prelude();
let root: Map<FfiString, i64> = [(FfiString::from("a"),
1i64)].into_iter().collect();
let error = match structural_walk(
&root,
@@ -733,7 +274,6 @@ fn handler_error_crosses_map_traversal() {
#[test]
fn interrupt_stops_without_running_remaining_callbacks() {
- test_prelude();
let root = Array::new(vec![1i64, 2, 3]);
let mut integers = 0;
let outcome = structural_walk(
@@ -785,7 +325,6 @@ impl StructuralVisitor for ManualRegionVisitor {
#[test]
fn manual_child_visit_can_override_def_region() {
- test_prelude();
let root = Array::new(vec![7i64, 8]);
let mut probe = ManualRegionVisitor::default();
assert!(structural_visit(&root, &mut probe).unwrap().is_none());
@@ -809,7 +348,6 @@ impl GeneratedLeafVisitor {
#[test]
fn generated_visitor_defaults_unmatched_values() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let mut visitor = GeneratedLeafVisitor::default();
assert!(structural_visit(&root, &mut visitor).unwrap().is_none());
@@ -851,7 +389,6 @@ impl GeneratedRecursiveVisitor {
#[test]
fn generated_visitor_can_drive_recursion_through_mut_self() {
- test_prelude();
let root = Array::new(vec![1i64, 2, 3]);
let mut visitor = GeneratedRecursiveVisitor::default();
let interrupt = structural_visit(&root, &mut visitor).unwrap().unwrap();
@@ -888,7 +425,6 @@ impl GenericDispatchProbe {
#[test]
fn generated_dispatch_supports_pod_and_ordered_catch_all() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let mut probe = GenericDispatchProbe::default();
assert!(structural_walk(&root, &mut probe, WalkOrder::PreOrder)
@@ -934,7 +470,6 @@ impl StructuralVisitor for StraddleVisitor {
#[test]
fn visitor_can_straddle_default_children() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let mut probe = StraddleVisitor::default();
assert!(structural_visit(&root, &mut probe).unwrap().is_none());
@@ -971,7 +506,6 @@ impl OrderProbe {
#[test]
fn stateful_structural_walk_supports_post_order() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let mut probe = OrderProbe::default();
assert!(structural_walk(&root, &mut probe, WalkOrder::PostOrder)
@@ -982,7 +516,6 @@ fn stateful_structural_walk_supports_post_order() {
#[test]
fn nested_walk_restores_the_outer_active_visitor() {
- test_prelude();
let outer = Array::new(vec![10i64, 20]);
let inner = Array::new(vec![1i64, 2]);
let mut entered_inner = false;
@@ -1020,7 +553,6 @@ fn nested_walk_restores_the_outer_active_visitor() {
#[test]
fn interrupt_payload_is_returned_to_the_caller() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let outcome = structural_walk(
&root,
@@ -1041,7 +573,6 @@ fn interrupt_payload_is_returned_to_the_caller() {
#[test]
fn handler_errors_include_native_visit_path() {
- test_prelude();
let root = Array::new(vec![1i64]);
let error = match structural_walk(
&root,
@@ -1063,7 +594,6 @@ fn handler_errors_include_native_visit_path() {
#[test]
fn visitor_errors_include_native_visit_path() {
- test_prelude();
struct FailingVisitor;
impl StructuralVisitor for FailingVisitor {
@@ -1102,7 +632,6 @@ fn visitor_errors_include_native_visit_path() {
#[test]
fn callback_panics_resume_after_the_registered_hook_returns() {
- test_prelude();
let root = Array::new(vec![1i64]);
let panic = match std::panic::catch_unwind(std::panic::AssertUnwindSafe(||
{
structural_walk(
@@ -1123,7 +652,6 @@ fn
callback_panics_resume_after_the_registered_hook_returns() {
#[test]
fn visitor_interrupt_propagates_through_default_children() {
- test_prelude();
struct InterruptingVisitor;
impl StructuralVisitor for InterruptingVisitor {
@@ -1151,7 +679,6 @@ fn visitor_interrupt_propagates_through_default_children()
{
#[test]
fn closure_walk_receives_def_region_kind() {
- test_prelude();
// C++: StructuralWalk<kPreOrder>(root,
// [&](const TVarObj* var, TVMFFIDefRegionKind kind) { ... })
let root = Array::new(vec![1i64, 2]);
@@ -1173,7 +700,6 @@ fn closure_walk_receives_def_region_kind() {
#[test]
fn closure_walk_supports_post_order_and_skip() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let mut order_probe = Vec::new();
assert!(structural_walk(
@@ -1213,7 +739,6 @@ fn closure_walk_supports_post_order_and_skip() {
// ---------------------------------------------------------------------------
#[test]
fn chain_accepts_owned_object_ref_links() {
- test_prelude();
let root = Array::new(vec![Array::new(vec![1i64]), Array::new(vec![2i64,
3])]);
let mut lengths = Vec::new();
assert!(structural_walk(
@@ -1236,7 +761,6 @@ fn chain_accepts_owned_object_ref_links() {
#[test]
fn chain_links_may_mix_def_region_arity() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let mut kinds = Vec::new();
let mut objects = 0;
@@ -1263,7 +787,6 @@ fn chain_links_may_mix_def_region_arity() {
#[test]
fn chain_links_can_skip_children() {
- test_prelude();
let root = Array::new(vec![Array::new(vec![1i64]),
Array::new(vec![2i64])]);
let mut arrays = 0;
let mut integers = 0;
@@ -1289,7 +812,6 @@ fn chain_links_can_skip_children() {
#[test]
fn chain_link_errors_include_native_visit_path() {
- test_prelude();
let root = Array::new(vec![1i64]);
let error = match structural_walk(
&root,
@@ -1308,7 +830,6 @@ fn chain_link_errors_include_native_visit_path() {
#[test]
fn chain_supports_post_order() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let events = std::cell::RefCell::new(Vec::new());
assert!(structural_walk(
@@ -1345,7 +866,6 @@ impl ObjectCounter {
#[test]
fn chain_splices_dispatch_walkers_between_closures() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let mut counter = ObjectCounter::default();
let mut integers = 0;
@@ -1365,7 +885,6 @@ fn chain_splices_dispatch_walkers_between_closures() {
#[test]
fn chain_supports_full_arity() {
- test_prelude();
let root = Array::new(vec![1i64, 2, 3]);
let mut integers = Vec::new();
let mut objects = 0;
@@ -1406,7 +925,6 @@ fn chain_supports_full_arity() {
#[test]
fn typed_lambda_walks_bare_and_as_single_link_tuple() {
- test_prelude();
// 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.
@@ -1438,7 +956,6 @@ fn typed_lambda_walks_bare_and_as_single_link_tuple() {
#[test]
fn bare_node_lambda_takes_def_region_kind() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let mut objects = 0;
assert!(structural_walk(
@@ -1479,7 +996,6 @@ impl StructuralVisitor for InheritedRegionProbe {
#[test]
fn def_region_is_inherited_through_containers() {
- test_prelude();
let root = Array::new(vec![Array::new(vec![1i64, 2])]);
let mut probe = InheritedRegionProbe {
at_root: true,
@@ -1490,116 +1006,55 @@ fn def_region_is_inherited_through_containers() {
}
#[test]
-fn reflected_field_def_region_reaches_typed_handler() {
- test_prelude();
- let root = RustVisitDefRegion {
- data: ObjectArc::new(RustVisitDefRegionObj {
- base: Object::new(),
- recursive: Any::from(1i64),
- plain: Any::from(2i64),
- non_recursive: Any::from(3i64),
- both: Any::from(4i64),
- ignored: Any::from(5i64),
- }),
- };
- let mut seen = Vec::new();
+fn reflected_fields_reach_typed_handlers() {
+ // Reference the existing test library so its C++ startup registrations
are linked.
+ assert_eq!(
+ unsafe { tvm_ffi::tvm_ffi_sys::TVMFFITestingDummyTarget() },
+ 0
+ );
+ let root = Function::get_global("ffi.MakeObjectFromPackedArgs")
+ .unwrap()
+ .call_tuple((
+ FfiString::from("testing.TestObjectBase"),
+ FfiString::from("v_i64"),
+ 1i64,
+ FfiString::from("v_f64"),
+ 2.5f64,
+ FfiString::from("v_str"),
+ FfiString::from("a reflected string"),
+ ))
+ .unwrap();
+ let mut integers = Vec::new();
+ let mut floats = Vec::new();
+ let mut strings = Vec::new();
assert!(structural_walk(
&root,
- |_value: i64, kind: DefRegionKind| {
- seen.push(kind);
- WalkResult::Advance
- },
+ (
+ |value: i64, kind: DefRegionKind| {
+ integers.push((value, kind));
+ WalkResult::Advance
+ },
+ |value: f64| {
+ floats.push(value);
+ WalkResult::Advance
+ },
+ |value: FfiString| {
+ strings.push(value);
+ WalkResult::Advance
+ },
+ ),
WalkOrder::PreOrder,
)
.unwrap()
.is_none());
- assert_eq!(
- seen,
- vec![
- DefRegionKind::Recursive,
- DefRegionKind::None,
- DefRegionKind::NonRecursive,
- DefRegionKind::NonRecursive,
- ]
- );
-}
-
-struct FreeVarClampProbe {
- at_root: bool,
- seen: Vec<(&'static str, DefRegionKind)>,
-}
-
-impl StructuralVisitor for FreeVarClampProbe {
- fn visit(
- &mut self,
- value: &VisitValue,
- def_region_kind: DefRegionKind,
- ) -> Result<Option<VisitInterrupt>> {
- if self.at_root {
- self.at_root = false;
- let root = value.cast::<Array<ObjectRef>>().unwrap();
- for index in 0..root.len() {
- let child = root.get(index).unwrap();
- if let Some(interrupt) = self.visit_child(&child,
DefRegionKind::NonRecursive)? {
- return Ok(Some(interrupt));
- }
- }
- return Ok(None);
- }
-
- if value.as_node::<RustVisitDefRegionObj>().is_some() {
- self.seen.push(("free_var", def_region_kind));
- } else if value.cast::<Array<i64>>().is_some() {
- self.seen.push(("array", def_region_kind));
- } else if let Some(integer) = value.cast::<i64>() {
- self.seen.push((
- if integer == 6 {
- "free_child"
- } else {
- "array_child"
- },
- def_region_kind,
- ));
- }
- self.default_visit_children(value, def_region_kind)
- }
-}
-
-#[test]
-fn non_recursive_region_is_clamped_for_free_var_children_only() {
- test_prelude();
- let free_var = RustVisitDefRegion {
- data: ObjectArc::new(RustVisitDefRegionObj {
- base: Object::new(),
- recursive: Any::new(),
- plain: Any::from(6i64),
- non_recursive: Any::new(),
- both: Any::new(),
- ignored: Any::new(),
- }),
- };
- let free_var: ObjectRef = free_var.try_cast().unwrap();
- let array: ObjectRef = Array::new(vec![7i64]).try_cast().unwrap();
- let root = Array::new(vec![free_var, array]);
- let mut probe = FreeVarClampProbe {
- at_root: true,
- seen: Vec::new(),
- };
- assert!(structural_visit(&root, &mut probe).unwrap().is_none());
- assert_eq!(
- probe.seen,
- vec![
- ("free_var", DefRegionKind::NonRecursive),
- ("free_child", DefRegionKind::None),
- ("array", DefRegionKind::NonRecursive),
- ("array_child", DefRegionKind::NonRecursive),
- ]
- );
+ assert_eq!(integers, vec![(1, DefRegionKind::None)]);
+ assert_eq!(floats, vec![2.5]);
+ assert_eq!(strings.len(), 1);
+ assert_eq!(strings[0].as_str(), "a reflected string");
}
#[test]
fn nested_tuple_chain_exceeds_flat_arity() {
- test_prelude();
let root = Array::new(vec![1i64, 2, 3]);
let mut integers = Vec::new();
let mut objects = 0;
@@ -1650,7 +1105,6 @@ fn nested_tuple_chain_exceeds_flat_arity() {
#[test]
fn nested_tuple_first_match_order_is_flattened() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let mut first = 0;
let mut second = 0;
@@ -1676,7 +1130,6 @@ fn nested_tuple_first_match_order_is_flattened() {
#[test]
fn callback_visit_defaults_only_when_no_link_matches() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let integers = Cell::new(0);
assert!(
@@ -1724,7 +1177,6 @@ fn stateful_visit_integer(value: i64, visitor: &mut
VisitContext<'_, StatefulVis
#[test]
fn stateful_callback_visit_uses_ordinary_mutable_state() {
- test_prelude();
let root = Array::new(vec![1i64, 2, 3]);
let mut visitor = VisitCallbacks::new(
StatefulVisitStats::default(),
@@ -1749,7 +1201,6 @@ struct StatefulVisitDepth {
#[test]
fn stateful_callback_visit_reborrows_visitor_during_recursion() {
- test_prelude();
let root = Array::new(vec![Array::new(vec![1i64, 2])]);
let mut visitor = VisitCallbacks::new(
StatefulVisitDepth::default(),
@@ -1773,7 +1224,6 @@ fn
stateful_callback_visit_reborrows_visitor_during_recursion() {
#[test]
fn callback_visit_can_reenter_the_same_fn_through_visitor() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let visits = Cell::new(0);
assert!(structural_visit(
@@ -1790,7 +1240,6 @@ fn
callback_visit_can_reenter_the_same_fn_through_visitor() {
#[test]
fn callback_visit_tuple_is_first_match_and_can_interrupt() {
- test_prelude();
let root = Array::new(vec![1i64, 2, 3]);
let fallback = Cell::new(0);
let interrupted = structural_visit(
@@ -1812,26 +1261,16 @@ fn
callback_visit_tuple_is_first_match_and_can_interrupt() {
}
#[test]
-fn callback_visit_supports_node_links_nested_tuples_and_def_regions() {
- test_prelude();
- let root = RustVisitDefRegion {
- data: ObjectArc::new(RustVisitDefRegionObj {
- base: Object::new(),
- recursive: Any::from(1i64),
- plain: Any::from(2i64),
- non_recursive: Any::from(3i64),
- both: Any::from(4i64),
- ignored: Any::from(5i64),
- }),
- };
+fn callback_visit_supports_node_links_and_nested_tuples() {
+ let root = Array::new(vec![1i64, 2]);
let seen = RefCell::new(Vec::new());
-
assert!(structural_visit(
&root,
(
(
|_value: f64, _visitor: &mut VisitContext<'_, ()>| {},
- |_node: &RustVisitDefRegionObj, visitor: &mut VisitContext<'_,
()>| {
+ |_node: &tvm_ffi::collections::array::ArrayObj,
+ visitor: &mut VisitContext<'_, ()>| {
assert_eq!(visitor.def_region_kind(), DefRegionKind::None);
visitor.visit_children()
},
@@ -1845,18 +1284,12 @@ fn
callback_visit_supports_node_links_nested_tuples_and_def_regions() {
.is_none());
assert_eq!(
*seen.borrow(),
- vec![
- (1, DefRegionKind::Recursive),
- (2, DefRegionKind::None),
- (3, DefRegionKind::NonRecursive),
- (4, DefRegionKind::NonRecursive),
- ]
+ vec![(1, DefRegionKind::None), (2, DefRegionKind::None)]
);
}
#[test]
fn callback_visit_with_overrides_child_def_region() {
- test_prelude();
let root = Array::new(vec![1i64, 2]);
let seen = RefCell::new(Vec::new());
assert!(structural_visit(
@@ -1885,7 +1318,6 @@ fn callback_visit_with_overrides_child_def_region() {
#[test]
fn nested_callback_visit_restores_the_outer_active_visitor() {
- test_prelude();
let outer = Array::new(vec![10i64, 20]);
let inner = Array::new(vec![1i64, 2]);
let entered_inner = Cell::new(false);
@@ -1914,7 +1346,6 @@ fn
nested_callback_visit_restores_the_outer_active_visitor() {
#[test]
fn callback_visit_panics_resume_and_leave_the_next_run_usable() {
- test_prelude();
let root = Array::new(vec![1i64]);
let entered_callback = Cell::new(false);
let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {