use std::sync::Arc;
use cljrs_ir::osr::build_osr_function;
use cljrs_ir::{Block, BlockId, Const, Inst, IrFunction, KnownFn, Terminator, VarId};
use cljrs_runtime::tiered::{Env, ir_interp::interpret_ir};
use cljrs_value::Value;
fn sum_loop_fn() -> IrFunction {
let mut f = IrFunction::new(Some(Arc::from("sum-to")), None);
let v = |n: u32| VarId(n);
f.params = vec![(Arc::from("n"), v(0))];
f.next_var = 9;
f.next_block = 4;
f.blocks = vec![
Block {
id: BlockId(0),
phis: vec![],
insts: vec![
Inst::Const(v(1), Const::Long(0)),
Inst::Const(v(2), Const::Long(0)),
],
terminator: Terminator::Jump(BlockId(1)),
},
Block {
id: BlockId(1),
phis: vec![
Inst::Phi(v(3), vec![(BlockId(0), v(1)), (BlockId(2), v(6))]),
Inst::Phi(v(4), vec![(BlockId(0), v(2)), (BlockId(2), v(7))]),
],
insts: vec![Inst::CallKnown(v(5), KnownFn::Lt, vec![v(3), v(0)])],
terminator: Terminator::Branch {
cond: v(5),
then_block: BlockId(2),
else_block: BlockId(3),
},
},
Block {
id: BlockId(2),
phis: vec![],
insts: vec![
Inst::Const(v(8), Const::Long(1)),
Inst::CallKnown(v(6), KnownFn::Add, vec![v(3), v(8)]),
Inst::CallKnown(v(7), KnownFn::Add, vec![v(4), v(3)]),
],
terminator: Terminator::RecurJump {
target: BlockId(1),
args: vec![v(6), v(7)],
},
},
Block {
id: BlockId(3),
phis: vec![],
insts: vec![],
terminator: Terminator::Return(v(4)),
},
];
f
}
fn run(ir: &IrFunction, args: Vec<Value>) -> Value {
let _mutator = cljrs_gc::register_mutator();
let globals = cljrs_runtime::Runtime::builder()
.execution_mode(cljrs_runtime::ExecutionMode::TreeWalk)
.build()
.expect("runtime")
.into_globals();
let mut env = Env::new(globals.clone(), "user");
let ns: Arc<str> = Arc::from("user");
cljrs_runtime::env::callback::push_eval_context(&env);
let result = interpret_ir(ir, args, &globals, &ns, &mut env);
cljrs_runtime::env::callback::pop_eval_context();
result.expect("interpret")
}
#[test]
fn osr_variant_resumes_mid_loop_state() {
let orig = sum_loop_fn();
assert_eq!(run(&orig, vec![Value::Long(10)]), Value::Long(45));
let osr = build_osr_function(&orig, BlockId(1)).expect("transform");
assert_eq!(osr.live_ins, vec![VarId(3), VarId(4), VarId(0)]);
let resumed = run(
&osr.func,
vec![Value::Long(5), Value::Long(10), Value::Long(10)],
);
assert_eq!(resumed, Value::Long(45));
}
#[test]
fn osr_variant_matches_original_for_every_entry_point() {
let orig = sum_loop_fn();
let osr = build_osr_function(&orig, BlockId(1)).expect("transform");
let n = 12i64;
let expected = run(&orig, vec![Value::Long(n)]);
for k in 0..=n {
let acc: i64 = (0..k).sum();
let resumed = run(
&osr.func,
vec![Value::Long(k), Value::Long(acc), Value::Long(n)],
);
assert_eq!(resumed, expected, "diverged when resuming at i={k}");
}
}
#[test]
fn osr_variant_runs_post_loop_code() {
let mut orig = sum_loop_fn();
let result_var = VarId(9);
orig.next_var = 10;
orig.blocks[3] = Block {
id: BlockId(3),
phis: vec![],
insts: vec![Inst::CallKnown(
result_var,
KnownFn::Add,
vec![VarId(4), VarId(0)],
)],
terminator: Terminator::Return(result_var),
};
assert_eq!(run(&orig, vec![Value::Long(10)]), Value::Long(55));
let osr = build_osr_function(&orig, BlockId(1)).expect("transform");
let resumed = run(
&osr.func,
vec![Value::Long(5), Value::Long(10), Value::Long(10)],
);
assert_eq!(resumed, Value::Long(55));
}