Skip to main content

trustfall_core/interpreter/
trace.rs

1use std::{
2    cell::RefCell, collections::BTreeMap, fmt::Debug, marker::PhantomData, num::NonZeroUsize,
3    rc::Rc, sync::Arc,
4};
5
6use serde::{de::DeserializeOwned, Deserialize, Serialize};
7
8use crate::{
9    interpreter::{Adapter, DataContext},
10    ir::{EdgeParameters, Eid, FieldValue, IRQuery, Vid},
11    util::BTreeMapTryInsertExt,
12};
13
14use super::{
15    AsVertex, ContextIterator, ContextOutcomeIterator, ResolveEdgeInfo, ResolveInfo, VertexInfo,
16    VertexIterator,
17};
18
19#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
20pub struct Opid(pub NonZeroUsize); // operation ID
21
22#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
23#[serde(bound = "Vertex: Debug + Clone + Serialize + DeserializeOwned")]
24pub struct Trace<Vertex> {
25    pub ops: BTreeMap<Opid, TraceOp<Vertex>>,
26
27    pub ir_query: IRQuery,
28
29    #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
30    pub arguments: BTreeMap<String, FieldValue>,
31}
32
33impl<Vertex> Trace<Vertex>
34where
35    Vertex: Clone + Debug + PartialEq + Eq + Serialize + DeserializeOwned,
36{
37    #[allow(dead_code)]
38    pub fn new(ir_query: IRQuery, arguments: BTreeMap<String, FieldValue>) -> Self {
39        Self { ops: Default::default(), ir_query, arguments }
40    }
41
42    pub fn record(&mut self, content: TraceOpContent<Vertex>, parent: Option<Opid>) -> Opid {
43        let next_opid = Opid(NonZeroUsize::new(self.ops.len() + 1).unwrap());
44
45        let op = TraceOp { opid: next_opid, parent_opid: parent, content };
46        self.ops.insert_or_error(next_opid, op).unwrap();
47        next_opid
48    }
49}
50
51#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
52#[serde(bound = "Vertex: Debug + Clone + Serialize + DeserializeOwned")]
53pub struct TraceOp<Vertex> {
54    pub opid: Opid,
55    pub parent_opid: Option<Opid>, // None parent_opid means this is a top-level operation
56
57    pub content: TraceOpContent<Vertex>,
58}
59
60#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
61#[serde(bound = "Vertex: Debug + Clone + Serialize + DeserializeOwned")]
62pub enum TraceOpContent<Vertex> {
63    // TODO: make a way to differentiate between different queries recorded in the same trace
64    Call(FunctionCall),
65
66    AdvanceInputIterator,
67    YieldInto(DataContext<Vertex>),
68    YieldFrom(YieldValue<Vertex>),
69
70    InputIteratorExhausted,
71    OutputIteratorExhausted,
72
73    ProduceQueryResult(BTreeMap<Arc<str>, FieldValue>),
74}
75
76#[allow(clippy::enum_variant_names)] // the variant names match the functions they represent
77#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
78pub enum FunctionCall {
79    ResolveStartingVertices(Vid),             // vertex ID
80    ResolveProperty(Vid, Arc<str>, Arc<str>), // vertex ID + type name + name of the property
81    ResolveNeighbors(Vid, Arc<str>, Eid),     // vertex ID + type name + edge ID
82    ResolveCoercion(Vid, Arc<str>, Arc<str>), // vertex ID + current type + coerced-to type
83}
84
85#[allow(clippy::enum_variant_names)] // the variant names match the functions they represent
86#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
87#[serde(bound = "Vertex: Debug + Clone + Serialize + DeserializeOwned")]
88pub enum YieldValue<Vertex> {
89    ResolveStartingVertices(Vertex),
90    ResolveProperty(DataContext<Vertex>, FieldValue),
91    ResolveNeighborsOuter(DataContext<Vertex>),
92    ResolveNeighborsInner(usize, Vertex), // iterable index + produced element
93    ResolveCoercion(DataContext<Vertex>, bool),
94}
95
96pub struct OnIterEnd<T, I: Iterator<Item = T>, F: FnOnce()> {
97    inner: I,
98    on_end_func: Option<F>,
99}
100
101impl<T, I, F> Iterator for OnIterEnd<T, I, F>
102where
103    F: FnOnce(),
104    I: Iterator<Item = T>,
105{
106    type Item = T;
107
108    fn next(&mut self) -> Option<Self::Item> {
109        let result = self.inner.next();
110        if result.is_none() {
111            let end_func = self.on_end_func.take();
112            if let Some(func) = end_func {
113                func();
114            }
115        }
116        result
117    }
118}
119
120fn make_iter_with_end_action<T, I: Iterator<Item = T>, F: FnOnce()>(
121    inner: I,
122    on_end: F,
123) -> OnIterEnd<T, I, F> {
124    OnIterEnd { inner, on_end_func: Some(on_end) }
125}
126
127pub struct PreActionIter<T, I: Iterator<Item = T>, F: Fn()> {
128    inner: I,
129    pre_action: F,
130}
131
132impl<T, I, F> Iterator for PreActionIter<T, I, F>
133where
134    F: Fn(),
135    I: Iterator<Item = T>,
136{
137    type Item = T;
138
139    fn next(&mut self) -> Option<Self::Item> {
140        (self.pre_action)();
141        self.inner.next()
142    }
143}
144
145fn make_iter_with_pre_action<T, I: Iterator<Item = T>, F: Fn()>(
146    inner: I,
147    pre_action: F,
148) -> PreActionIter<T, I, F> {
149    PreActionIter { inner, pre_action }
150}
151
152/// An adapter "middleware" that records all adapter operations into a linear, replayable trace.
153///
154/// Tapping adapters must be done at top level, on the top-level adapter which is being used
155/// for query execution i.e. where the `<V>` generic on the resolver methods
156/// is the same as `AdapterT::Vertex`.
157///
158/// Otherwise, the recorded traces may not be possible to replay since they would be incomplete:
159/// they would only capture a portion of the execution, the rest of which is missing.
160#[derive(Debug, Clone)]
161pub struct AdapterTap<'vertex, AdapterT>
162where
163    AdapterT: Adapter<'vertex>,
164    AdapterT::Vertex: Clone + Debug + PartialEq + Eq + Serialize + DeserializeOwned + 'vertex,
165{
166    tracer: Rc<RefCell<Trace<AdapterT::Vertex>>>,
167    inner: AdapterT,
168    _phantom: PhantomData<&'vertex ()>,
169}
170
171impl<'vertex, AdapterT> AdapterTap<'vertex, AdapterT>
172where
173    AdapterT: Adapter<'vertex>,
174    AdapterT::Vertex: Clone + Debug + PartialEq + Eq + Serialize + DeserializeOwned + 'vertex,
175{
176    pub fn new(adapter: AdapterT, tracer: Rc<RefCell<Trace<AdapterT::Vertex>>>) -> Self {
177        Self { tracer, inner: adapter, _phantom: PhantomData }
178    }
179
180    pub fn finish(self) -> Trace<AdapterT::Vertex> {
181        // Ensure nothing is reading the trace i.e. we can safely stop interpreting.
182        let trace_ref = self.tracer.borrow_mut();
183        let new_trace = Trace::new(trace_ref.ir_query.clone(), trace_ref.arguments.clone());
184        drop(trace_ref);
185        self.tracer.replace(new_trace)
186    }
187}
188
189pub fn tap_results<'vertex, AdapterT>(
190    adapter_tap: Arc<AdapterTap<'vertex, AdapterT>>,
191    result_iter: impl Iterator<Item = BTreeMap<Arc<str>, FieldValue>> + 'vertex,
192) -> impl Iterator<Item = BTreeMap<Arc<str>, FieldValue>> + 'vertex
193where
194    AdapterT: Adapter<'vertex> + 'vertex,
195    AdapterT::Vertex: Clone + Debug + PartialEq + Eq + Serialize + DeserializeOwned + 'vertex,
196{
197    result_iter.inspect(move |result| {
198        adapter_tap
199            .tracer
200            .borrow_mut()
201            .record(TraceOpContent::ProduceQueryResult(result.clone()), None);
202    })
203}
204
205impl<'vertex, AdapterT> Adapter<'vertex> for AdapterTap<'vertex, AdapterT>
206where
207    AdapterT: Adapter<'vertex>,
208    AdapterT::Vertex: Clone + Debug + PartialEq + Eq + Serialize + DeserializeOwned + 'vertex,
209{
210    type Vertex = AdapterT::Vertex;
211
212    fn resolve_starting_vertices(
213        &self,
214        edge_name: &Arc<str>,
215        parameters: &EdgeParameters,
216        resolve_info: &ResolveInfo,
217    ) -> VertexIterator<'vertex, Self::Vertex> {
218        let mut trace = self.tracer.borrow_mut();
219        let call_opid = trace.record(
220            TraceOpContent::Call(FunctionCall::ResolveStartingVertices(resolve_info.vid())),
221            None,
222        );
223        drop(trace);
224
225        let inner_iter = self.inner.resolve_starting_vertices(edge_name, parameters, resolve_info);
226        let tracer_ref_1 = self.tracer.clone();
227        let tracer_ref_2 = self.tracer.clone();
228        Box::new(
229            make_iter_with_end_action(inner_iter, move || {
230                tracer_ref_1
231                    .borrow_mut()
232                    .record(TraceOpContent::OutputIteratorExhausted, Some(call_opid));
233            })
234            .inspect(move |vertex| {
235                tracer_ref_2.borrow_mut().record(
236                    TraceOpContent::YieldFrom(YieldValue::ResolveStartingVertices(vertex.clone())),
237                    Some(call_opid),
238                );
239            }),
240        )
241    }
242
243    fn resolve_property<V: AsVertex<Self::Vertex> + 'vertex>(
244        &self,
245        contexts: ContextIterator<'vertex, V>,
246        type_name: &Arc<str>,
247        property_name: &Arc<str>,
248        resolve_info: &ResolveInfo,
249    ) -> ContextOutcomeIterator<'vertex, V, FieldValue> {
250        let mut trace = self.tracer.borrow_mut();
251        let call_opid = trace.record(
252            TraceOpContent::Call(FunctionCall::ResolveProperty(
253                resolve_info.vid(),
254                type_name.clone(),
255                property_name.clone(),
256            )),
257            None,
258        );
259        drop(trace);
260
261        let tracer_ref_1 = self.tracer.clone();
262        let tracer_ref_2 = self.tracer.clone();
263        let tracer_ref_3 = self.tracer.clone();
264        let wrapped_contexts = Box::new(
265            make_iter_with_end_action(
266                make_iter_with_pre_action(contexts, move || {
267                    tracer_ref_1
268                        .borrow_mut()
269                        .record(TraceOpContent::AdvanceInputIterator, Some(call_opid));
270                }),
271                move || {
272                    tracer_ref_2
273                        .borrow_mut()
274                        .record(TraceOpContent::InputIteratorExhausted, Some(call_opid));
275                },
276            )
277            .inspect(move |context| {
278                tracer_ref_3.borrow_mut().record(
279                    TraceOpContent::YieldInto(context.clone().flat_map(&mut |v| v.into_vertex())),
280                    Some(call_opid),
281                );
282            }),
283        );
284        let inner_iter =
285            self.inner.resolve_property(wrapped_contexts, type_name, property_name, resolve_info);
286
287        let tracer_ref_4 = self.tracer.clone();
288        let tracer_ref_5 = self.tracer.clone();
289        Box::new(
290            make_iter_with_end_action(inner_iter, move || {
291                tracer_ref_4
292                    .borrow_mut()
293                    .record(TraceOpContent::OutputIteratorExhausted, Some(call_opid));
294            })
295            .map(move |(context, value)| {
296                tracer_ref_5.borrow_mut().record(
297                    TraceOpContent::YieldFrom(YieldValue::ResolveProperty(
298                        context.clone().flat_map(&mut |v| v.into_vertex()),
299                        value.clone(),
300                    )),
301                    Some(call_opid),
302                );
303
304                (context, value)
305            }),
306        )
307    }
308
309    fn resolve_neighbors<V: AsVertex<Self::Vertex> + 'vertex>(
310        &self,
311        contexts: ContextIterator<'vertex, V>,
312        type_name: &Arc<str>,
313        edge_name: &Arc<str>,
314        parameters: &EdgeParameters,
315        resolve_info: &ResolveEdgeInfo,
316    ) -> ContextOutcomeIterator<'vertex, V, VertexIterator<'vertex, Self::Vertex>> {
317        let mut trace = self.tracer.borrow_mut();
318        let call_opid = trace.record(
319            TraceOpContent::Call(FunctionCall::ResolveNeighbors(
320                resolve_info.origin_vid(),
321                type_name.clone(),
322                resolve_info.eid(),
323            )),
324            None,
325        );
326        drop(trace);
327
328        let tracer_ref_1 = self.tracer.clone();
329        let tracer_ref_2 = self.tracer.clone();
330        let tracer_ref_3 = self.tracer.clone();
331        let wrapped_contexts = Box::new(
332            make_iter_with_end_action(
333                make_iter_with_pre_action(contexts, move || {
334                    tracer_ref_1
335                        .borrow_mut()
336                        .record(TraceOpContent::AdvanceInputIterator, Some(call_opid));
337                }),
338                move || {
339                    tracer_ref_2
340                        .borrow_mut()
341                        .record(TraceOpContent::InputIteratorExhausted, Some(call_opid));
342                },
343            )
344            .inspect(move |context| {
345                tracer_ref_3.borrow_mut().record(
346                    TraceOpContent::YieldInto(context.clone().flat_map(&mut |v| v.into_vertex())),
347                    Some(call_opid),
348                );
349            }),
350        );
351        let inner_iter = self.inner.resolve_neighbors(
352            wrapped_contexts,
353            type_name,
354            edge_name,
355            parameters,
356            resolve_info,
357        );
358
359        let tracer_ref_4 = self.tracer.clone();
360        let tracer_ref_5 = self.tracer.clone();
361        Box::new(
362            make_iter_with_end_action(inner_iter, move || {
363                tracer_ref_4
364                    .borrow_mut()
365                    .record(TraceOpContent::OutputIteratorExhausted, Some(call_opid));
366            })
367            .map(move |(context, neighbor_iter)| {
368                let mut trace = tracer_ref_5.borrow_mut();
369                let outer_iterator_opid = trace.record(
370                    TraceOpContent::YieldFrom(YieldValue::ResolveNeighborsOuter(
371                        context.clone().flat_map(&mut |v| v.into_vertex()),
372                    )),
373                    Some(call_opid),
374                );
375                drop(trace);
376
377                let tracer_ref_6 = tracer_ref_5.clone();
378                let tapped_neighbor_iter = neighbor_iter.enumerate().map(move |(pos, vertex)| {
379                    tracer_ref_6.borrow_mut().record(
380                        TraceOpContent::YieldFrom(YieldValue::ResolveNeighborsInner(
381                            pos,
382                            vertex.clone(),
383                        )),
384                        Some(outer_iterator_opid),
385                    );
386
387                    vertex
388                });
389
390                let tracer_ref_7 = tracer_ref_5.clone();
391                let final_neighbor_iter: VertexIterator<'vertex, Self::Vertex> =
392                    Box::new(make_iter_with_end_action(tapped_neighbor_iter, move || {
393                        tracer_ref_7.borrow_mut().record(
394                            TraceOpContent::OutputIteratorExhausted,
395                            Some(outer_iterator_opid),
396                        );
397                    }));
398
399                (context, final_neighbor_iter)
400            }),
401        )
402    }
403
404    fn resolve_coercion<V: AsVertex<Self::Vertex> + 'vertex>(
405        &self,
406        contexts: ContextIterator<'vertex, V>,
407        type_name: &Arc<str>,
408        coerce_to_type: &Arc<str>,
409        resolve_info: &ResolveInfo,
410    ) -> ContextOutcomeIterator<'vertex, V, bool> {
411        let mut trace = self.tracer.borrow_mut();
412        let call_opid = trace.record(
413            TraceOpContent::Call(FunctionCall::ResolveCoercion(
414                resolve_info.vid(),
415                type_name.clone(),
416                coerce_to_type.clone(),
417            )),
418            None,
419        );
420        drop(trace);
421
422        let tracer_ref_1 = self.tracer.clone();
423        let tracer_ref_2 = self.tracer.clone();
424        let tracer_ref_3 = self.tracer.clone();
425        let wrapped_contexts = Box::new(
426            make_iter_with_end_action(
427                make_iter_with_pre_action(contexts, move || {
428                    tracer_ref_1
429                        .borrow_mut()
430                        .record(TraceOpContent::AdvanceInputIterator, Some(call_opid));
431                }),
432                move || {
433                    tracer_ref_2
434                        .borrow_mut()
435                        .record(TraceOpContent::InputIteratorExhausted, Some(call_opid));
436                },
437            )
438            .inspect(move |context| {
439                tracer_ref_3.borrow_mut().record(
440                    TraceOpContent::YieldInto(context.clone().flat_map(&mut |v| v.into_vertex())),
441                    Some(call_opid),
442                );
443            }),
444        );
445        let inner_iter =
446            self.inner.resolve_coercion(wrapped_contexts, type_name, coerce_to_type, resolve_info);
447
448        let tracer_ref_4 = self.tracer.clone();
449        let tracer_ref_5 = self.tracer.clone();
450        Box::new(
451            make_iter_with_end_action(inner_iter, move || {
452                tracer_ref_4
453                    .borrow_mut()
454                    .record(TraceOpContent::OutputIteratorExhausted, Some(call_opid));
455            })
456            .map(move |(context, can_coerce)| {
457                tracer_ref_5.borrow_mut().record(
458                    TraceOpContent::YieldFrom(YieldValue::ResolveCoercion(
459                        context.clone().flat_map(&mut |v| v.into_vertex()),
460                        can_coerce,
461                    )),
462                    Some(call_opid),
463                );
464
465                (context, can_coerce)
466            }),
467        )
468    }
469}