Skip to main content

hara_native/jit/
backend.rs

1use super::trace_ir::{
2    ExitReason, ExitSnapshot, NumericVectorSlice, Trace, TraceOp, TraceOutcome, TraceValue,
3};
4use crate::core::{IntrinsicOp, Value};
5
6pub trait TraceBackend {
7    type Compiled;
8    fn compile(&mut self, trace: &Trace) -> Result<Self::Compiled, String>;
9    fn enter(
10        &mut self,
11        trace: &mut Self::Compiled,
12        locals: &mut [TraceValue],
13        max_iterations: u32,
14    ) -> TraceOutcome;
15}
16
17#[derive(Default)]
18pub struct CheckedBackend;
19
20impl TraceBackend for CheckedBackend {
21    type Compiled = Trace;
22
23    fn compile(&mut self, trace: &Trace) -> Result<Trace, String> {
24        Ok(trace.clone())
25    }
26
27    fn enter(
28        &mut self,
29        trace: &mut Trace,
30        locals: &mut [TraceValue],
31        max_iterations: u32,
32    ) -> TraceOutcome {
33        for operation in &trace.operations {
34            let valid = match operation {
35                TraceOp::GuardLocalI64 { local } => {
36                    matches!(locals.get(usize::from(*local)), Some(TraceValue::I64(_)))
37                }
38                TraceOp::GuardLocalBool { local } => {
39                    matches!(locals.get(usize::from(*local)), Some(TraceValue::Bool(_)))
40                }
41                TraceOp::GuardLocalNil { local } => {
42                    matches!(locals.get(usize::from(*local)), Some(TraceValue::Nil))
43                }
44                TraceOp::GuardLocalVectorI64 { local } => {
45                    locals.get(usize::from(*local)).is_some_and(numeric_vector)
46                }
47                _ => true,
48            };
49            if !valid {
50                return TraceOutcome::SideExit {
51                    reason: ExitReason::WrongTag,
52                    iterations: 0,
53                    snapshot: ExitSnapshot {
54                        function: trace.function,
55                        instruction: trace.resume_ip,
56                        locals: locals.to_vec(),
57                        stack: Vec::new(),
58                    },
59                };
60            }
61        }
62        let mut iterations = 0;
63        let mut stack = Vec::with_capacity(8);
64        while iterations < max_iterations {
65            let checkpoint = locals.to_vec();
66            stack.clear();
67            for operation in &trace.operations {
68                let exit = |reason| TraceOutcome::SideExit {
69                    reason,
70                    iterations,
71                    snapshot: ExitSnapshot {
72                        function: trace.function,
73                        instruction: trace.resume_ip,
74                        locals: checkpoint.clone(),
75                        stack: Vec::new(),
76                    },
77                };
78                match *operation {
79                    TraceOp::GuardLocalI64 { .. }
80                    | TraceOp::GuardLocalBool { .. }
81                    | TraceOp::GuardLocalNil { .. }
82                    | TraceOp::GuardLocalVectorI64 { .. } => {}
83                    TraceOp::LoadLocal { local } => match locals.get(local as usize).cloned() {
84                        Some(value) => stack.push(value),
85                        None => return exit(ExitReason::WrongTag),
86                    },
87                    TraceOp::ConstantI64(value) => stack.push(TraceValue::I64(value)),
88                    TraceOp::ConstantBool(value) => stack.push(TraceValue::Bool(value)),
89                    TraceOp::ConstantNil => stack.push(TraceValue::Nil),
90                    TraceOp::ConstantVectorI64 { vector } => {
91                        let Some(values) = trace.vectors.get(usize::from(vector)) else {
92                            return exit(ExitReason::Unsupported);
93                        };
94                        stack.push(TraceValue::Indexed(Box::new(Value::Vector(
95                            values.iter().copied().map(Value::Number).collect(),
96                        ))));
97                    }
98                    TraceOp::StoreLocal { local } => {
99                        let Some(value) = stack.pop() else {
100                            return exit(ExitReason::Unsupported);
101                        };
102                        let Some(slot) = locals.get_mut(local as usize) else {
103                            return exit(ExitReason::Unsupported);
104                        };
105                        *slot = value;
106                    }
107                    TraceOp::Pop => {
108                        stack.pop();
109                    }
110                    TraceOp::GuardTruthy { expected } => {
111                        let Some(value) = stack.pop() else {
112                            return exit(ExitReason::Unsupported);
113                        };
114                        let truthy = !matches!(value, TraceValue::Bool(false) | TraceValue::Nil);
115                        if truthy != expected {
116                            return exit(ExitReason::BranchChanged);
117                        }
118                    }
119                    TraceOp::BinaryI64(op) => {
120                        let (Some(TraceValue::I64(right)), Some(TraceValue::I64(left))) =
121                            (stack.pop(), stack.pop())
122                        else {
123                            return exit(ExitReason::WrongTag);
124                        };
125                        let value = match op {
126                            IntrinsicOp::Add => left.checked_add(right).map(TraceValue::I64),
127                            IntrinsicOp::Subtract => left.checked_sub(right).map(TraceValue::I64),
128                            IntrinsicOp::Multiply => left.checked_mul(right).map(TraceValue::I64),
129                            IntrinsicOp::Divide | IntrinsicOp::Remainder | IntrinsicOp::Modulo
130                                if right == 0 =>
131                            {
132                                return exit(ExitReason::DivisionByZero)
133                            }
134                            IntrinsicOp::Divide => left.checked_div(right).map(TraceValue::I64),
135                            IntrinsicOp::Remainder | IntrinsicOp::Modulo => {
136                                if left == i64::MIN && right == -1 {
137                                    Some(TraceValue::I64(0))
138                                } else {
139                                    let Some(remainder) = left.checked_rem(right) else {
140                                        return exit(ExitReason::Overflow);
141                                    };
142                                    Some(TraceValue::I64(remainder))
143                                }
144                            }
145                            IntrinsicOp::Less => Some(TraceValue::Bool(left < right)),
146                            IntrinsicOp::LessOrEqual => Some(TraceValue::Bool(left <= right)),
147                            IntrinsicOp::Greater => Some(TraceValue::Bool(left > right)),
148                            IntrinsicOp::GreaterOrEqual => Some(TraceValue::Bool(left >= right)),
149                            IntrinsicOp::Equal => Some(TraceValue::Bool(left == right)),
150                        };
151                        let Some(value) = value else {
152                            return exit(ExitReason::Overflow);
153                        };
154                        stack.push(value);
155                    }
156                    TraceOp::VectorCountI64 => {
157                        let Some(vector) =
158                            stack.pop().and_then(|value| numeric_vector_values(&value))
159                        else {
160                            return exit(ExitReason::WrongTag);
161                        };
162                        stack.push(TraceValue::I64(vector.len() as i64));
163                    }
164                    TraceOp::VectorFirstI64 | TraceOp::VectorSecondI64 => {
165                        let index = usize::from(matches!(*operation, TraceOp::VectorSecondI64));
166                        let Some(vector) =
167                            stack.pop().and_then(|value| numeric_vector_values(&value))
168                        else {
169                            return exit(ExitReason::WrongTag);
170                        };
171                        let Some(value) = vector.get(index).copied() else {
172                            return exit(ExitReason::IndexOutOfBounds);
173                        };
174                        stack.push(TraceValue::I64(value));
175                    }
176                    TraceOp::VectorRestI64 => {
177                        let Some(vector) =
178                            stack.pop().and_then(|value| numeric_vector_values(&value))
179                        else {
180                            return exit(ExitReason::WrongTag);
181                        };
182                        stack.push(TraceValue::VectorSlice(Box::new(NumericVectorSlice {
183                            start: usize::from(!vector.is_empty()),
184                            values: vector,
185                        })));
186                    }
187                    TraceOp::VectorNthI64 => {
188                        let Some(TraceValue::I64(index)) = stack.pop() else {
189                            return exit(ExitReason::WrongTag);
190                        };
191                        let Some(vector) = stack.pop() else {
192                            return exit(ExitReason::WrongTag);
193                        };
194                        let Some(index) = usize::try_from(index).ok() else {
195                            return exit(ExitReason::IndexOutOfBounds);
196                        };
197                        let Some(values) = numeric_vector_values(&vector) else {
198                            return exit(ExitReason::WrongTag);
199                        };
200                        let Some(value) = values.get(index).copied() else {
201                            return exit(ExitReason::IndexOutOfBounds);
202                        };
203                        stack.push(TraceValue::I64(value));
204                    }
205                    TraceOp::LoopBackedge => iterations += 1,
206                }
207            }
208        }
209        TraceOutcome::Completed { iterations }
210    }
211}
212
213fn numeric_vector(value: &TraceValue) -> bool {
214    numeric_vector_values(value).is_some()
215}
216
217fn numeric_vector_values(value: &TraceValue) -> Option<Vec<i64>> {
218    match value {
219        TraceValue::Indexed(value) => match value.as_ref() {
220            Value::Tuple(values) => values
221                .iter()
222                .map(|value| match value {
223                    Value::Number(value) => Some(*value),
224                    _ => None,
225                })
226                .collect(),
227            Value::Vector(values) => values
228                .iter()
229                .map(|value| match value {
230                    Value::Number(value) => Some(*value),
231                    _ => None,
232                })
233                .collect(),
234            _ => None,
235        },
236        TraceValue::VectorSlice(slice) => Some(slice.values[slice.start..].to_vec()),
237        _ => None,
238    }
239}
240
241#[cfg(test)]
242mod tests {
243    use super::*;
244
245    fn increment_trace() -> Trace {
246        Trace {
247            function: 0,
248            header: 2,
249            resume_ip: 2,
250            operations: vec![
251                TraceOp::GuardLocalI64 { local: 0 },
252                TraceOp::LoadLocal { local: 0 },
253                TraceOp::ConstantI64(1),
254                TraceOp::BinaryI64(IntrinsicOp::Add),
255                TraceOp::StoreLocal { local: 0 },
256                TraceOp::LoopBackedge,
257            ],
258            vectors: Vec::new(),
259        }
260    }
261
262    #[test]
263    fn checked_backend_executes_and_guards() {
264        let trace = increment_trace();
265        let mut backend = CheckedBackend;
266        let mut compiled = backend.compile(&trace).unwrap();
267        let mut locals = [TraceValue::I64(0)];
268        assert_eq!(
269            backend.enter(&mut compiled, &mut locals, 5),
270            TraceOutcome::Completed { iterations: 5 }
271        );
272        assert_eq!(locals[0], TraceValue::I64(5));
273        let mut wrong = [TraceValue::Bool(false)];
274        assert!(matches!(
275            backend.enter(&mut compiled, &mut wrong, 1),
276            TraceOutcome::SideExit {
277                reason: ExitReason::WrongTag,
278                ..
279            }
280        ));
281    }
282
283    #[test]
284    fn checked_backend_indexes_numeric_vector_constants_and_exits_on_bounds() {
285        let trace = Trace {
286            function: 0,
287            header: 0,
288            resume_ip: 0,
289            operations: vec![
290                TraceOp::ConstantVectorI64 { vector: 0 },
291                TraceOp::LoadLocal { local: 0 },
292                TraceOp::VectorNthI64,
293                TraceOp::StoreLocal { local: 1 },
294                TraceOp::LoopBackedge,
295            ],
296            vectors: vec![vec![10, 20, 30]],
297        };
298        let mut backend = CheckedBackend;
299        let mut compiled = backend.compile(&trace).unwrap();
300        let mut locals = [TraceValue::I64(1), TraceValue::Nil];
301        assert_eq!(
302            backend.enter(&mut compiled, &mut locals, 1),
303            TraceOutcome::Completed { iterations: 1 }
304        );
305        assert_eq!(locals[1], TraceValue::I64(20));
306
307        locals[0] = TraceValue::I64(3);
308        assert!(matches!(
309            backend.enter(&mut compiled, &mut locals, 1),
310            TraceOutcome::SideExit {
311                reason: ExitReason::IndexOutOfBounds,
312                ..
313            }
314        ));
315    }
316
317    #[test]
318    fn checked_backend_restores_the_iteration_checkpoint_on_failure() {
319        let trace = Trace {
320            function: 0,
321            header: 4,
322            resume_ip: 4,
323            operations: vec![
324                TraceOp::LoadLocal { local: 0 },
325                TraceOp::ConstantI64(1),
326                TraceOp::BinaryI64(IntrinsicOp::Add),
327                TraceOp::StoreLocal { local: 0 },
328                TraceOp::LoadLocal { local: 0 },
329                TraceOp::ConstantI64(0),
330                TraceOp::BinaryI64(IntrinsicOp::Divide),
331                TraceOp::Pop,
332                TraceOp::LoopBackedge,
333            ],
334            vectors: Vec::new(),
335        };
336        let mut compiled = trace.clone();
337        let mut locals = [TraceValue::I64(9)];
338        assert!(matches!(
339            CheckedBackend.enter(&mut compiled, &mut locals, 1),
340            TraceOutcome::SideExit {
341                reason: ExitReason::DivisionByZero,
342                snapshot: ExitSnapshot { locals, .. },
343                ..
344            } if locals == vec![TraceValue::I64(9)]
345        ));
346    }
347}