tlopex commented on code in PR #704: URL: https://github.com/apache/tvm-ffi/pull/704#discussion_r3763354516
########## rust/tvm-ffi/src/extra/structural_mutate.rs: ########## @@ -0,0 +1,1609 @@ +/* + * 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 mapping. +//! +//! [`structural_map`] mirrors the callback-controlled C++ `StructuralMap`, +//! but keeps recursion, typed dispatch, definition-region propagation, and +//! identity remapping in Rust. The root is consumed so Rust ownership and +//! the runtime strong count jointly define the boundary for optional +//! in-place container mutation. Passing a clone naturally selects +//! copy-on-write behavior. +//! +//! Rust owns callback dispatch, memoization, identity remapping, and most +//! container traversal. Map/Dict storage is traversed through a narrow C ABI +//! that calls directly back into Rust for each value, allowing unique maps to +//! be updated in place without exposing the runtime's private hash layout. +//! A non-container object with a foreign `__s_mutate__` or +//! `__s_maybe_inplace_mutate__` hook is rejected rather than silently +//! replacing its custom semantics with reflection. + +use std::collections::HashMap; +use std::ffi::c_void; +use std::marker::PhantomData; +use std::ops::{ControlFlow, Deref}; +use std::ptr::NonNull; +use std::sync::atomic::{AtomicUsize, Ordering}; + +use crate::any::{Any, AnyView}; +use crate::error::{Error, Result, RUNTIME_ERROR, TYPE_ERROR}; +use crate::function::Function; +use crate::object::{self, ObjectCore}; +use crate::tvm_ffi_sys::TVMFFIFieldFlagBitMask::{ + kTVMFFIFieldFlagBitMaskSEqHashIgnore, kTVMFFIFieldFlagBitSetterIsFunctionObj, +}; +use crate::tvm_ffi_sys::{ + TVMFFIAny, TVMFFIFieldInfo, TVMFFIFieldSetter, TVMFFIFunctionCall, TVMFFIGetTypeInfo, + TVMFFIMapMutateValues, TVMFFIObject, TVMFFITypeAttrColumn, TVMFFITypeIndex, +}; +use crate::tvm_ffi_sys::{TVMFFIObjectHandle, TVMFFISEqHashKind}; + +use super::structural_visit::{ + field_def_region, for_each_field, free_var_child_region, type_attr_column, type_key_of, + DefRegionKind, SeqPrefix, TypeAttrColumn, VisitValue, WalkOrder, +}; + +const STRUCTURAL_MUTATE_ATTR: &str = "__s_mutate__"; +const STRUCTURAL_MAYBE_INPLACE_MUTATE_ATTR: &str = "__s_maybe_inplace_mutate__"; +const SHALLOW_COPY_ATTR: &str = "__ffi_shallow_copy__"; +const FLAG_SEQ_HASH_IGNORE: i64 = kTVMFFIFieldFlagBitMaskSEqHashIgnore as i64; +const FLAG_SETTER_IS_FUNCTION: i64 = kTVMFFIFieldFlagBitSetterIsFunctionObj as i64; + +/// Borrowed value passed to structural-map callbacks. +/// +/// This is the same representation used by structural-visit callbacks; the +/// alias keeps typed casts and borrowed node checks on one audited unsafe +/// implementation. +pub type MapValue = VisitValue; + +/// Result type produced by a structural-map callback. +#[doc(hidden)] +pub type MapResult = Result<Any>; + +/// Convert an infallible or fallible callback result into [`MapResult`]. +pub trait IntoMapResult { + fn into_map_result(self) -> MapResult; +} + +impl IntoMapResult for Any { + #[inline] + fn into_map_result(self) -> MapResult { + Ok(self) + } +} + +impl IntoMapResult for Result<Any> { + #[inline] + fn into_map_result(self) -> MapResult { + self + } +} + +/// Ordered typed replacement dispatch for [`structural_map`]. +/// +/// `None` means no handler matched and preserves the current value. A +/// generated `#[dispatch(map)]` implementation tests `map_*` methods in +/// source order and returns the first match. +pub trait MapDispatch: Sized { + fn dispatch_map( + &mut self, + value: &MapValue, + def_region_kind: DefRegionKind, + ) -> Option<MapResult>; +} + +impl<V: MapDispatch> MapDispatch for &mut V { + #[inline] + fn dispatch_map( + &mut self, + value: &MapValue, + def_region_kind: DefRegionKind, + ) -> Option<MapResult> { + (**self).dispatch_map(value, def_region_kind) + } +} + +/// Conversion into the mapper consumed by [`structural_map`]. +#[diagnostic::on_unimplemented( + message = "unsupported structural-map callback shape", + label = "this value cannot be used as a structural mapper", + note = "pass `&mut` a type implementing `MapDispatch`, a supported closure, or a tuple of callbacks" +)] +pub trait IntoMapper<Marker> { + type Mapper: MapDispatch; + fn into_mapper(self) -> Self::Mapper; +} + +#[doc(hidden)] +pub enum ByMapDispatch {} + +impl<'a, V: MapDispatch> IntoMapper<ByMapDispatch> for &'a mut V { + type Mapper = &'a mut V; + + #[inline] + fn into_mapper(self) -> Self::Mapper { + self + } +} + +/// One typed callback in a structural-map tuple. +/// +/// Links are tried in tuple order and the first matching link supplies the +/// replacement. Supported shapes are owned FFI values, borrowed object +/// nodes, and `&MapValue`, each optionally followed by [`DefRegionKind`]. +/// Numeric links match the complete FFI `Int` or `Float` type tag and then use +/// Rust `as` conversion semantics, so prefer `i64` and `f64` unless narrowing +/// is intentional. +pub trait MapChainLink<Marker>: sealed_map::SealedMapLink<Marker> { + #[doc(hidden)] + fn try_map(&mut self, value: &MapValue, def_region_kind: DefRegionKind) -> Option<MapResult>; +} + +mod sealed_map { + use super::{DefRegionKind, IntoMapResult, MapDispatch, MapValue, ObjectCore}; + + pub trait SealedMapLink<Marker> {} + + impl<F, T, O> SealedMapLink<super::ByMapOwned<T>> for F + where + F: FnMut(T) -> O, + O: IntoMapResult, + { + } + + impl<F, T, O> SealedMapLink<super::ByMapOwnedKind<T>> for F + where + F: FnMut(T, DefRegionKind) -> O, + O: IntoMapResult, + { + } + + impl<F, N: ObjectCore, O> SealedMapLink<super::ByMapNode<N>> for F + where + F: for<'a> FnMut(&'a N) -> O, + O: IntoMapResult, + { + } + + impl<F, N: ObjectCore, O> SealedMapLink<super::ByMapNodeKind<N>> for F + where + F: for<'a> FnMut(&'a N, DefRegionKind) -> O, + O: IntoMapResult, + { + } + + impl<F, O> SealedMapLink<super::ByMapCatchAll> for F + where + F: for<'a> FnMut(&'a MapValue) -> O, + O: IntoMapResult, + { + } + + impl<F, O> SealedMapLink<super::ByMapCatchAllKind> for F + where + F: for<'a> FnMut(&'a MapValue, DefRegionKind) -> O, + O: IntoMapResult, + { + } + + impl<V: MapDispatch> SealedMapLink<super::ByMapDispatchLink> for &mut V {} +} + +#[doc(hidden)] +pub struct ByMapOwned<T>(PhantomData<T>); + +impl<F, T, O> MapChainLink<ByMapOwned<T>> for F +where + F: FnMut(T) -> O, + T: crate::type_traits::AnyCompatible, + O: IntoMapResult, +{ + #[inline] + fn try_map(&mut self, value: &MapValue, _def_region_kind: DefRegionKind) -> Option<MapResult> { + value.cast::<T>().map(|typed| self(typed).into_map_result()) + } +} + +#[doc(hidden)] +pub struct ByMapOwnedKind<T>(PhantomData<T>); + +impl<F, T, O> MapChainLink<ByMapOwnedKind<T>> for F +where + F: FnMut(T, DefRegionKind) -> O, + T: crate::type_traits::AnyCompatible, + O: IntoMapResult, +{ + #[inline] + fn try_map(&mut self, value: &MapValue, def_region_kind: DefRegionKind) -> Option<MapResult> { + value + .cast::<T>() + .map(|typed| self(typed, def_region_kind).into_map_result()) + } +} + +#[doc(hidden)] +pub struct ByMapNode<N>(PhantomData<N>); + +impl<F, N, O> MapChainLink<ByMapNode<N>> for F +where + F: for<'a> FnMut(&'a N) -> O, + N: ObjectCore, + O: IntoMapResult, +{ + #[inline] + fn try_map(&mut self, value: &MapValue, _def_region_kind: DefRegionKind) -> Option<MapResult> { + value + .as_node::<N>() + .map(|node| self(node).into_map_result()) + } +} + +#[doc(hidden)] +pub struct ByMapNodeKind<N>(PhantomData<N>); + +impl<F, N, O> MapChainLink<ByMapNodeKind<N>> for F +where + F: for<'a> FnMut(&'a N, DefRegionKind) -> O, + N: ObjectCore, + O: IntoMapResult, +{ + #[inline] + fn try_map(&mut self, value: &MapValue, def_region_kind: DefRegionKind) -> Option<MapResult> { + value + .as_node::<N>() + .map(|node| self(node, def_region_kind).into_map_result()) + } +} + +#[doc(hidden)] +pub enum ByMapCatchAll {} + +impl<F, O> MapChainLink<ByMapCatchAll> for F +where + F: for<'a> FnMut(&'a MapValue) -> O, + O: IntoMapResult, +{ + #[inline] + fn try_map(&mut self, value: &MapValue, _def_region_kind: DefRegionKind) -> Option<MapResult> { + Some(self(value).into_map_result()) + } +} + +#[doc(hidden)] +pub enum ByMapCatchAllKind {} + +impl<F, O> MapChainLink<ByMapCatchAllKind> for F +where + F: for<'a> FnMut(&'a MapValue, DefRegionKind) -> O, + O: IntoMapResult, +{ + #[inline] + fn try_map(&mut self, value: &MapValue, def_region_kind: DefRegionKind) -> Option<MapResult> { + Some(self(value, def_region_kind).into_map_result()) + } +} + +#[doc(hidden)] +pub enum ByMapDispatchLink {} + +impl<V: MapDispatch> MapChainLink<ByMapDispatchLink> for &mut V { + #[inline] + fn try_map(&mut self, value: &MapValue, def_region_kind: DefRegionKind) -> Option<MapResult> { + self.dispatch_map(value, def_region_kind) + } +} + +/// Statically dispatched tuple mapper, public only as an [`IntoMapper`] +/// projection. +#[doc(hidden)] +pub struct MapChain<Links, Markers> { + links: Links, + markers: PhantomData<fn(Markers)>, +} + +macro_rules! impl_map_chain { + ($(($F:ident, $M:ident, $idx:tt)),+) => { + impl<$($F, $M,)+> MapDispatch for MapChain<($($F,)+), ($($M,)+)> + where + $($F: MapChainLink<$M>,)+ + { + #[inline] + fn dispatch_map( + &mut self, + value: &MapValue, + def_region_kind: DefRegionKind, + ) -> Option<MapResult> { + $( + if let Some(result) = self.links.$idx.try_map(value, def_region_kind) { + return Some(result); + } + )+ + None + } + } + + impl<$($F, $M,)+> IntoMapper<($($M,)+)> for ($($F,)+) + where + $($F: MapChainLink<$M>,)+ + { + type Mapper = MapChain<($($F,)+), ($($M,)+)>; + + #[inline] + fn into_mapper(self) -> Self::Mapper { + MapChain { + links: self, + markers: PhantomData, + } + } + } + }; +} + +impl_map_chain!((F0, M0, 0)); Review Comment: Rust requires separate implementations for tuples containing 1 to 8 callbacks. Already added comments -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected] --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
