use vyre_foundation::error::IrError;
use vyre_foundation::ir::{BufferAccess, BufferDecl, DataType, Expr, Node, Program};
use vyre_foundation::transform::inline::inline_calls_with_resolver;
fn callee_with_builtin(built_in_expr: Expr) -> Program {
Program::wrapped(
vec![
BufferDecl::storage("x", 0, BufferAccess::ReadOnly, DataType::U32).with_count(64),
BufferDecl::output("result", 1, DataType::U32).with_count(64),
],
[64, 1, 1],
vec![Node::Store {
buffer: "result".into(),
index: built_in_expr,
value: Expr::u32(42),
}],
)
}
fn caller_for(op_id: &str) -> Program {
Program::wrapped(
vec![
BufferDecl::storage("x", 0, BufferAccess::ReadOnly, DataType::U32).with_count(64),
BufferDecl::output("out", 1, DataType::U32).with_count(64),
],
[64, 1, 1],
vec![Node::Store {
buffer: "out".into(),
index: Expr::u32(0),
value: Expr::call(op_id, vec![Expr::load("x", Expr::u32(0))]),
}],
)
}
fn builtin_resolver(id: &str) -> Option<Program> {
let built_in = match id {
"uses_gid" => Expr::InvocationId { axis: 0 },
"uses_wgid" => Expr::WorkgroupId { axis: 0 },
"uses_lid" => Expr::LocalId { axis: 0 },
"uses_sgid" => Expr::SubgroupLocalId,
"uses_sgsize" => Expr::SubgroupSize,
"uses_gid_in_binop" => {
return Some(Program::wrapped(
vec![
BufferDecl::storage("a", 0, BufferAccess::ReadOnly, DataType::U32)
.with_count(64),
BufferDecl::output("out2", 1, DataType::U32).with_count(64),
],
[64, 1, 1],
vec![Node::Store {
buffer: "out2".into(),
index: Expr::add(Expr::InvocationId { axis: 1 }, Expr::u32(1)),
value: Expr::u32(7),
}],
));
}
_ => return None,
};
Some(callee_with_builtin(built_in))
}
fn assert_lowering_error_names_fix(result: Result<Program, IrError>, builtin_name: &str) {
match result {
Err(IrError::Lowering { message }) => {
assert!(
message.contains("Fix:"),
"Fix: inline rejection for {builtin_name} must include 'Fix:' guidance, \
got message: {message:?}"
);
let names_a_builtin = message.contains("InvocationId")
|| message.contains("WorkgroupId")
|| message.contains("LocalId")
|| message.contains("SubgroupLocalId")
|| message.contains("SubgroupSize");
assert!(
names_a_builtin,
"Fix: inline rejection for {builtin_name} must name the offending \
built-in class in the error message, got: {message:?}"
);
}
Ok(_) => panic!(
"Fix: inlining a callee that uses {builtin_name} must return Err(Lowering), \
not Ok, the silent replacement of per-invocation built-ins with 0 \
was a miscompile."
),
Err(other) => panic!(
"Fix: inlining a callee that uses {builtin_name} must return \
Err(Lowering), got unexpected error variant {other:?}"
),
}
}
#[test]
fn inline_rejects_callee_with_invocation_id() {
let caller = caller_for("uses_gid");
let result = inline_calls_with_resolver(&caller, builtin_resolver);
assert_lowering_error_names_fix(result, "InvocationId");
}
#[test]
fn inline_rejects_callee_with_workgroup_id() {
let caller = caller_for("uses_wgid");
let result = inline_calls_with_resolver(&caller, builtin_resolver);
assert_lowering_error_names_fix(result, "WorkgroupId");
}
#[test]
fn inline_rejects_callee_with_local_id() {
let caller = caller_for("uses_lid");
let result = inline_calls_with_resolver(&caller, builtin_resolver);
assert_lowering_error_names_fix(result, "LocalId");
}
#[test]
fn inline_rejects_callee_with_subgroup_local_id() {
let caller = caller_for("uses_sgid");
let result = inline_calls_with_resolver(&caller, builtin_resolver);
assert_lowering_error_names_fix(result, "SubgroupLocalId");
}
#[test]
fn inline_rejects_callee_with_subgroup_size() {
let caller = caller_for("uses_sgsize");
let result = inline_calls_with_resolver(&caller, builtin_resolver);
assert_lowering_error_names_fix(result, "SubgroupSize");
}
#[test]
fn inline_rejects_callee_with_invocation_id_in_binop() {
let caller = caller_for("uses_gid_in_binop");
let result = inline_calls_with_resolver(&caller, builtin_resolver);
assert_lowering_error_names_fix(result, "InvocationId (inside BinOp)");
}