use lex_ast as a;
use lex_bytecode::{compile_program, vm::Vm, Value};
use crate::handler::DefaultHandler;
use crate::policy::Policy;
use lex_types::TypeError;
pub fn evaluate_examples(stages: &[a::Stage]) -> Vec<TypeError> {
let helpers = synthesize_helpers(stages);
if helpers.cases.is_empty() {
return Vec::new();
}
let mut augmented: Vec<a::Stage> = stages.to_vec();
augmented.extend(helpers.helper_stages);
let bc = compile_program(&augmented);
let bc = std::sync::Arc::new(bc);
let mut out = Vec::new();
for case in &helpers.cases {
match run_case(&bc, case) {
CaseOutcome::Pass => {}
CaseOutcome::Mismatch { expected, got } => {
out.push(TypeError::ExampleMismatch {
at_node: "n_0".into(),
fn_name: case.fn_name.clone(),
case_index: case.case_index,
expected,
got,
});
}
CaseOutcome::RuntimeError(msg) => {
out.push(TypeError::ExampleMismatch {
at_node: "n_0".into(),
fn_name: case.fn_name.clone(),
case_index: case.case_index,
expected: "(declared value)".into(),
got: format!("runtime error: {msg}"),
});
}
}
}
out
}
struct Helpers {
helper_stages: Vec<a::Stage>,
cases: Vec<Case>,
}
struct Case {
fn_name: String,
case_index: usize,
arg_helpers: Vec<String>,
expected_helper: String,
}
fn synthesize_helpers(stages: &[a::Stage]) -> Helpers {
let mut helper_stages = Vec::new();
let mut cases = Vec::new();
for stage in stages {
let a::Stage::FnDecl(fd) = stage else { continue };
if fd.examples.is_empty() {
continue;
}
if !fd.type_params.is_empty() {
continue;
}
if !fd.effects.is_empty() {
continue;
}
for (k, ex) in fd.examples.iter().enumerate() {
let mut arg_helpers = Vec::with_capacity(ex.args.len());
for (i, arg) in ex.args.iter().enumerate() {
let helper_name = format!("__ex_{}_{}_arg_{}", fd.name, k, i);
helper_stages.push(zero_arg_helper(&helper_name, fd.params[i].ty.clone(), arg.clone()));
arg_helpers.push(helper_name);
}
let expected_helper = format!("__ex_{}_{}_expected", fd.name, k);
helper_stages.push(zero_arg_helper(
&expected_helper,
fd.return_type.clone(),
ex.expected.clone(),
));
cases.push(Case {
fn_name: fd.name.clone(),
case_index: k,
arg_helpers,
expected_helper,
});
}
}
Helpers { helper_stages, cases }
}
fn zero_arg_helper(name: &str, return_type: a::TypeExpr, body: a::CExpr) -> a::Stage {
a::Stage::FnDecl(a::FnDecl {
name: name.into(),
type_params: Vec::new(),
params: Vec::new(),
effects: Vec::new(),
effect_row_var: None,
return_type,
body,
examples: Vec::new(),
})
}
enum CaseOutcome {
Pass,
Mismatch { expected: String, got: String },
RuntimeError(String),
}
fn run_case(bc: &std::sync::Arc<lex_bytecode::Program>, case: &Case) -> CaseOutcome {
let mut arg_values: Vec<Value> = Vec::with_capacity(case.arg_helpers.len());
for helper in &case.arg_helpers {
match call_zero_arg(bc, helper) {
Ok(v) => arg_values.push(v),
Err(e) => return CaseOutcome::RuntimeError(format!("computing arg from `{helper}`: {e}")),
}
}
let expected = match call_zero_arg(bc, &case.expected_helper) {
Ok(v) => v,
Err(e) => return CaseOutcome::RuntimeError(format!("computing expected from `{}`: {e}", case.expected_helper)),
};
let got = match call_with_args(bc, &case.fn_name, arg_values) {
Ok(v) => v,
Err(e) => return CaseOutcome::RuntimeError(format!("calling `{}`: {e}", case.fn_name)),
};
if expected == got {
CaseOutcome::Pass
} else {
CaseOutcome::Mismatch {
expected: pretty_value(&expected),
got: pretty_value(&got),
}
}
}
fn call_zero_arg(bc: &std::sync::Arc<lex_bytecode::Program>, name: &str) -> Result<Value, String> {
call_with_args(bc, name, Vec::new())
}
fn call_with_args(
bc: &std::sync::Arc<lex_bytecode::Program>,
name: &str,
args: Vec<Value>,
) -> Result<Value, String> {
let handler = DefaultHandler::new(Policy::pure()).with_program(std::sync::Arc::clone(bc));
let mut vm = Vm::with_handler(bc, Box::new(handler));
vm.set_step_limit(1_000_000);
vm.call(name, args).map_err(|e| format!("{e:?}"))
}
fn pretty_value(v: &Value) -> String {
match v {
Value::Int(n) => n.to_string(),
Value::Float(f) => f.to_string(),
Value::Bool(b) => b.to_string(),
Value::Str(s) => format!("{s:?}"),
Value::Unit => "()".into(),
Value::List(xs) => format!(
"[{}]",
xs.iter().map(pretty_value).collect::<Vec<_>>().join(", ")
),
Value::Tuple(xs) => format!(
"({})",
xs.iter().map(pretty_value).collect::<Vec<_>>().join(", ")
),
Value::Variant { name, args } if args.is_empty() => name.clone(),
Value::Variant { name, args } => format!(
"{}({})",
name,
args.iter().map(pretty_value).collect::<Vec<_>>().join(", ")
),
Value::Record { fields: fs, .. } => format!(
"{{ {} }}",
fs.iter()
.map(|(k, v)| format!("{k}: {}", pretty_value(v)))
.collect::<Vec<_>>()
.join(", ")
),
other => format!("{other:?}"),
}
}