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}