Skip to main content

tvm_ffi/extra/structural_visit/
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//! Reusable customization of default structural descent.
21
22use super::*;
23
24/// A reusable policy for managing context around default visit and walk recursion.
25///
26/// `visit_children()` on this policy's context continues with the next policy
27/// (or the built-in hooks and reflected fields). `visit()` re-enters the full
28/// callback engine for a child. A tuple `(outer, inner)` composes two policies;
29/// tuples may nest. Policies receive `&self`; mutable pass data belongs in the
30/// context's state. Save and restore scoped state around descent.
31pub trait ContextPolicy<State> {
32    /// Customize default descent for the current value.
33    ///
34    /// Return interrupts explicitly, and restore any scoped state before
35    /// returning an interrupt or error. Walk invokes this between its pre- and
36    /// post-order callback positions; visit invokes it only on callback miss
37    /// or when a matched callback requests default descent.
38    fn default_visit(
39        &self,
40        value: &StructuralView,
41        visitor: &mut VisitContext<'_, State>,
42    ) -> Result<Option<VisitInterrupt>>;
43}
44
45/// Default descent through registered hooks or reflected structural fields.
46pub struct DefaultContextPolicy;
47
48impl<State> ContextPolicy<State> for DefaultContextPolicy {
49    fn default_visit(
50        &self,
51        _value: &StructuralView,
52        visitor: &mut VisitContext<'_, State>,
53    ) -> Result<Option<VisitInterrupt>> {
54        visitor.visit_children()
55    }
56}
57
58impl<State, Outer: ContextPolicy<State>, Inner: ContextPolicy<State>> ContextPolicy<State>
59    for (Outer, Inner)
60{
61    fn default_visit(
62        &self,
63        value: &StructuralView,
64        visitor: &mut VisitContext<'_, State>,
65    ) -> Result<Option<VisitInterrupt>> {
66        let kind = visitor.def_region_kind();
67        visit_with_policy(
68            &mut NextPolicy {
69                driver: &mut *visitor.driver,
70                policy: &self.1,
71            },
72            &self.0,
73            value,
74            kind,
75        )
76    }
77}
78
79#[inline(always)]
80pub(super) fn visit_with_policy<State>(
81    driver: &mut dyn VisitContextDriver<State>,
82    policy: &impl ContextPolicy<State>,
83    value: &StructuralView,
84    def_region_kind: DefRegionKind,
85) -> Result<Option<VisitInterrupt>> {
86    with_visit_region(
87        def_region_kind,
88        #[inline(always)]
89        |def_region_kind| {
90            policy.default_visit(
91                value,
92                &mut VisitContext {
93                    driver,
94                    current: StructuralView::from_raw(value.raw()),
95                    def_region_kind,
96                    _not_send_sync: PhantomData,
97                },
98            )
99        },
100    )
101}
102
103struct NextPolicy<'a, State, Policy> {
104    driver: &'a mut dyn VisitContextDriver<State>,
105    policy: &'a Policy,
106}
107
108impl<State, Policy: ContextPolicy<State>> VisitContextDriver<State>
109    for NextPolicy<'_, State, Policy>
110{
111    fn state(&self) -> &State {
112        self.driver.state()
113    }
114    fn state_mut(&mut self) -> &mut State {
115        self.driver.state_mut()
116    }
117    fn visit_raw(&mut self, raw: TVMFFIAny, kind: DefRegionKind) -> Result<Option<VisitInterrupt>> {
118        self.driver.visit_raw(raw, kind)
119    }
120    fn visit_children_raw(
121        &mut self,
122        raw: TVMFFIAny,
123        kind: DefRegionKind,
124    ) -> Result<Option<VisitInterrupt>> {
125        visit_with_policy(
126            self.driver,
127            self.policy,
128            &StructuralView::from_raw(raw),
129            kind,
130        )
131    }
132}
133
134pub(super) struct VisitDescent<'a, V> {
135    pub(super) visitor: &'a mut V,
136}
137
138impl<State, V: StructuralVisitor + VisitCallbackState<State>> VisitContextDriver<State>
139    for VisitDescent<'_, V>
140{
141    fn state(&self) -> &State {
142        self.visitor.callback_state()
143    }
144    fn state_mut(&mut self) -> &mut State {
145        self.visitor.callback_state_mut()
146    }
147    fn visit_raw(&mut self, raw: TVMFFIAny, kind: DefRegionKind) -> Result<Option<VisitInterrupt>> {
148        VisitContextDriver::visit_raw(self.visitor, raw, kind)
149    }
150    fn visit_children_raw(
151        &mut self,
152        raw: TVMFFIAny,
153        kind: DefRegionKind,
154    ) -> Result<Option<VisitInterrupt>> {
155        // Bypass the current policy; children still re-enter the complete visitor.
156        default_user_visit_children(self.visitor, &StructuralView::from_raw(raw), kind)
157    }
158}
159
160/// A walk dispatcher combined with a reusable default-recursion policy.
161///
162/// The dispatcher is also the state visible through the policy's context. Use
163/// `#[dispatch(walk)]` or implement [`WalkDispatch`] to define its callbacks.
164/// Pass this value or a mutable reference to [`structural_walk`].
165pub struct WalkWithContextPolicy<Walker, Policy> {
166    walker: Walker,
167    policy: Rc<Policy>,
168}
169
170impl<Walker: WalkDispatch, Policy: ContextPolicy<Walker>> WalkWithContextPolicy<Walker, Policy> {
171    /// Combine a dispatcher and a default-recursion policy.
172    pub fn new(walker: Walker, policy: Policy) -> Self {
173        Self {
174            walker,
175            policy: Rc::new(policy),
176        }
177    }
178
179    /// Access the dispatcher and its traversal state.
180    pub fn state(&self) -> &Walker {
181        &self.walker
182    }
183
184    /// Mutably access the dispatcher outside an active walk.
185    pub fn state_mut(&mut self) -> &mut Walker {
186        &mut self.walker
187    }
188
189    /// Recover the dispatcher and its state.
190    pub fn into_state(self) -> Walker {
191        self.walker
192    }
193}
194
195#[doc(hidden)]
196pub enum ByPolicyWalk {}
197
198impl<Walker: WalkDispatch, Policy: ContextPolicy<Walker>> IntoWalker<ByPolicyWalk>
199    for WalkWithContextPolicy<Walker, Policy>
200{
201    type Walker = Self;
202    fn into_walker(self) -> Self {
203        self
204    }
205}
206
207impl<Walker: WalkDispatch, Policy: ContextPolicy<Walker>> IntoWalker<ByPolicyWalk>
208    for &mut WalkWithContextPolicy<Walker, Policy>
209{
210    type Walker = Self;
211    fn into_walker(self) -> Self {
212        self
213    }
214}
215
216impl<Walker: WalkDispatch, Policy: ContextPolicy<Walker>> NativeVisit
217    for WalkWithContextPolicy<Walker, Policy>
218{
219    const CUSTOM_DESCENT: bool = true;
220
221    fn visit(&mut self, value: &StructuralView, kind: DefRegionKind) -> Result<WalkResult> {
222        self.walker
223            .dispatch_walk(value, kind)
224            .unwrap_or_else(|| Ok(WalkResult::Advance))
225    }
226
227    fn default_visit_children<const PRE_ORDER: bool>(
228        &mut self,
229        value: &StructuralView,
230        kind: DefRegionKind,
231    ) -> Result<Option<VisitInterrupt>> {
232        let policy = Rc::clone(&self.policy);
233        visit_with_policy(
234            &mut WalkDescent::<_, _, PRE_ORDER> { visitor: self },
235            &*policy,
236            value,
237            kind,
238        )
239    }
240}
241
242struct WalkDescent<'a, Walker, Policy, const PRE_ORDER: bool> {
243    visitor: &'a mut WalkWithContextPolicy<Walker, Policy>,
244}
245
246impl<Walker: WalkDispatch, Policy: ContextPolicy<Walker>, const PRE_ORDER: bool>
247    VisitContextDriver<Walker> for WalkDescent<'_, Walker, Policy, PRE_ORDER>
248{
249    fn state(&self) -> &Walker {
250        &self.visitor.walker
251    }
252    fn state_mut(&mut self) -> &mut Walker {
253        &mut self.visitor.walker
254    }
255    fn visit_raw(&mut self, raw: TVMFFIAny, kind: DefRegionKind) -> Result<Option<VisitInterrupt>> {
256        if raw.type_index == TVMFFITypeIndex::kTVMFFINone as i32 {
257            return Ok(None);
258        }
259        let active = active_structural_visitor()?;
260        let context = std::ptr::from_mut(&mut *self.visitor).cast::<c_void>();
261        finish(with_current_visitor_context(active, context, || {
262            call_visitor(active, raw, kind)
263        }))
264    }
265    fn visit_children_raw(
266        &mut self,
267        raw: TVMFFIAny,
268        kind: DefRegionKind,
269    ) -> Result<Option<VisitInterrupt>> {
270        default_walk_children::<_, PRE_ORDER>(self.visitor, &StructuralView::from_raw(raw), kind)
271    }
272}