Skip to main content

tvm_ffi/extra/structural_mutate/
policy.rs

1/*
2 * Licensed to the Apache Software Foundation (ASF) under one
3 * or more contributor license agreements.  See the NOTICE file
4 * distributed with this work for additional information
5 * regarding copyright ownership.  The ASF licenses this file
6 * to you under the Apache License, Version 2.0 (the
7 * "License"); you may not use this file except in compliance
8 * with the License.  You may obtain a copy of the License at
9 *
10 *   http://www.apache.org/licenses/LICENSE-2.0
11 *
12 * Unless required by applicable law or agreed to in writing,
13 * software distributed under the License is distributed on an
14 * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15 * KIND, either express or implied.  See the License for the
16 * specific language governing permissions and limitations
17 * under the License.
18 */
19
20//! Context management around default structural mutation.
21
22use super::*;
23
24/// Default-recursion policy for [`MutateCallbacks::with_policy`] and [`MapWithContextPolicy`].
25///
26/// `ctx.default_maybe_inplace_mutate_result(value)` continues to the next policy,
27/// then hooks or reflected fields; children re-enter callback dispatch.
28/// Policies share callback state and compose as `(outer, inner)`.
29/// Restore user state before returning, including on errors; use
30/// [`MutateContext::with_def_region_kind`] for scoped definition regions.
31pub trait MutContextPolicy<State> {
32    /// Customize default descent, preserving the input permission and result marker.
33    fn default_mutate(
34        &self,
35        value: MutateValue<'_>,
36        ctx: &mut MutateContext<'_, State>,
37    ) -> Result<UnchangedOr<Any>>;
38}
39
40/// Default descent through registered hooks or reflected structural fields.
41pub struct DefaultMutContextPolicy;
42
43impl<State> MutContextPolicy<State> for DefaultMutContextPolicy {
44    fn default_mutate(
45        &self,
46        value: MutateValue<'_>,
47        ctx: &mut MutateContext<'_, State>,
48    ) -> Result<UnchangedOr<Any>> {
49        ctx.default_maybe_inplace_mutate_result(value)
50    }
51}
52
53impl<State, Outer: MutContextPolicy<State>, Inner: MutContextPolicy<State>> MutContextPolicy<State>
54    for (Outer, Inner)
55{
56    fn default_mutate(
57        &self,
58        value: MutateValue<'_>,
59        ctx: &mut MutateContext<'_, State>,
60    ) -> Result<UnchangedOr<Any>> {
61        let kind = ctx.def_region_kind();
62        mutate_with_policy(
63            &mut NextPolicy {
64                driver: &mut *ctx.driver,
65                policy: &self.1,
66            },
67            &self.0,
68            value,
69            kind,
70        )
71        .and_then(UnchangedOr::from_carrier)
72    }
73}
74
75#[inline(always)]
76pub(super) fn mutate_with_policy<State>(
77    driver: &mut dyn MutateContextDriver<State>,
78    policy: &impl MutContextPolicy<State>,
79    value: MutateValue<'_>,
80    kind: DefRegionKind,
81) -> Result<Any> {
82    let raw = value.value.raw();
83    with_mutation_region(
84        kind,
85        #[inline(always)]
86        |kind| {
87            let mut ctx = MutateContext {
88                driver,
89                def_region_kind: kind,
90                inplace_mode: value.inplace_mode(),
91                _not_send_sync: PhantomData,
92            };
93            policy.default_mutate(value, &mut ctx).map(Any::from)
94        },
95    )
96    .map_err(|error| with_value_context(error, raw))
97}
98
99struct NextPolicy<'a, State, Policy> {
100    driver: &'a mut dyn MutateContextDriver<State>,
101    policy: &'a Policy,
102}
103
104impl<State, Policy: MutContextPolicy<State>> MutateContextDriver<State>
105    for NextPolicy<'_, State, Policy>
106{
107    fn state(&self) -> &State {
108        self.driver.state()
109    }
110    fn state_mut(&mut self) -> &mut State {
111        self.driver.state_mut()
112    }
113    fn mutate_borrowed(&mut self, value: AnyView<'_>, kind: DefRegionKind) -> Result<Any> {
114        self.driver.mutate_borrowed(value, kind)
115    }
116    fn mutate_owned(&mut self, value: Any, kind: DefRegionKind, mode: InplaceMode) -> Result<Any> {
117        self.driver.mutate_owned(value, kind, mode)
118    }
119    fn default_mutate_borrowed(&mut self, value: AnyView<'_>, kind: DefRegionKind) -> Result<Any> {
120        let view = StructuralView::from_raw(*value.as_raw_ffi_any());
121        mutate_with_policy(self.driver, self.policy, MutateValue::borrowed(&view), kind)
122    }
123    fn default_mutate_value(
124        &mut self,
125        mut value: MutateValue<'_>,
126        kind: DefRegionKind,
127        mode: InplaceMode,
128    ) -> Result<Any> {
129        value.mode = value.permit(mode).inplace_mode(value.value.raw());
130        mutate_with_policy(self.driver, self.policy, value, kind)
131    }
132    fn var_remap_get(&mut self, var: &StructuralView) -> Result<Option<Any>> {
133        self.driver.var_remap_get(var)
134    }
135    fn var_remap_set(&mut self, var: &StructuralView, replacement: &Any) -> Result<()> {
136        self.driver.var_remap_set(var, replacement)
137    }
138}
139
140pub(super) struct MutationDescent<'a, Driver> {
141    pub(super) driver: &'a mut Driver,
142}
143
144impl<State, Driver: MutationDriver + MutateCallbackState<State>> MutateContextDriver<State>
145    for MutationDescent<'_, Driver>
146{
147    fn state(&self) -> &State {
148        self.driver.callback_state()
149    }
150    fn state_mut(&mut self) -> &mut State {
151        self.driver.callback_state_mut()
152    }
153    fn mutate_borrowed(&mut self, value: AnyView<'_>, kind: DefRegionKind) -> Result<Any> {
154        self.driver
155            .dispatch_raw(*value.as_raw_ffi_any(), kind, Permit::Copy)
156    }
157    fn mutate_owned(&mut self, value: Any, kind: DefRegionKind, mode: InplaceMode) -> Result<Any> {
158        let raw = *value.as_raw_ffi_any();
159        let result = self.driver.dispatch_raw(raw, kind, mode.permit())?;
160        Ok(if is_unchanged(&result) { value } else { result })
161    }
162    fn default_mutate_borrowed(&mut self, value: AnyView<'_>, kind: DefRegionKind) -> Result<Any> {
163        // The policy context has already installed this definition region.
164        default_mutate_driver(self.driver, *value.as_raw_ffi_any(), kind, Permit::Copy)
165    }
166    fn default_mutate_value(
167        &mut self,
168        value: MutateValue<'_>,
169        kind: DefRegionKind,
170        mode: InplaceMode,
171    ) -> Result<Any> {
172        let permit = value.permit(mode);
173        default_mutate_driver(self.driver, value.value.raw(), kind, permit)
174    }
175    fn var_remap_get(&mut self, var: &StructuralView) -> Result<Option<Any>> {
176        self.driver.var_remap_get_raw(var.raw())
177    }
178    fn var_remap_set(&mut self, var: &StructuralView, replacement: &Any) -> Result<()> {
179        self.driver.var_remap_set_raw(var.raw(), replacement)
180    }
181}
182
183impl<D, Policy, const PRE_ORDER: bool> MutateCallbackState<D>
184    for NativeMapper<'_, D, Policy, PRE_ORDER>
185{
186    fn callback_state(&self) -> &D {
187        self.dispatch
188    }
189    fn callback_state_mut(&mut self) -> &mut D {
190        self.dispatch
191    }
192}
193
194/// A [`MapDispatch`] with a [`MutContextPolicy`], sharing the dispatcher's state.
195///
196/// Pass this value or a mutable reference to [`structural_map`]. Each node's
197/// map callback runs outside its policy scope; its children run inside.
198/// This mapper cannot be a callback tuple member or another wrapper's dispatcher.
199/// Compose policies as `(outer, inner)` within one wrapper.
200pub struct MapWithContextPolicy<Mapper, Policy> {
201    mapper: Mapper,
202    policy: Rc<Policy>,
203}
204
205impl<Mapper: MapDispatch, Policy: MutContextPolicy<Mapper>> MapWithContextPolicy<Mapper, Policy> {
206    /// Combine a dispatcher and a default-recursion policy.
207    pub fn new(mapper: Mapper, policy: Policy) -> Self {
208        Self {
209            mapper,
210            policy: Rc::new(policy),
211        }
212    }
213
214    /// Access the dispatcher and shared state.
215    pub fn state(&self) -> &Mapper {
216        &self.mapper
217    }
218
219    /// Mutably access the dispatcher outside recursive calls.
220    pub fn state_mut(&mut self) -> &mut Mapper {
221        &mut self.mapper
222    }
223
224    /// Recover the dispatcher and its state.
225    pub fn into_state(self) -> Mapper {
226        self.mapper
227    }
228}
229
230impl<Mapper: MapDispatch, Policy: MutContextPolicy<Mapper>> NativeMap
231    for MapWithContextPolicy<Mapper, Policy>
232{
233    fn map_root(&mut self, root: Any, order: WalkOrder) -> Result<Any> {
234        run_native_mapper(root, &mut self.mapper, Some(self.policy.clone()), order)
235    }
236}
237
238impl<Mapper: MapDispatch, Policy: MutContextPolicy<Mapper>> NativeMap
239    for &mut MapWithContextPolicy<Mapper, Policy>
240{
241    fn map_root(&mut self, root: Any, order: WalkOrder) -> Result<Any> {
242        (**self).map_root(root, order)
243    }
244}
245
246#[doc(hidden)]
247pub enum ByPolicyMap {}
248
249impl<Mapper: MapDispatch, Policy: MutContextPolicy<Mapper>> IntoMapper<ByPolicyMap>
250    for MapWithContextPolicy<Mapper, Policy>
251{
252    type Mapper = Self;
253    fn into_mapper(self) -> Self {
254        self
255    }
256}
257
258impl<'a, Mapper: MapDispatch, Policy: MutContextPolicy<Mapper>> IntoMapper<ByPolicyMap>
259    for &'a mut MapWithContextPolicy<Mapper, Policy>
260{
261    type Mapper = Self;
262    fn into_mapper(self) -> Self {
263        self
264    }
265}