Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions src/values.rs
Original file line number Diff line number Diff line change
Expand Up @@ -327,8 +327,18 @@ impl Value {
native_assert!(*tag < info.variants.len(), "Variant index out of range.");

let payload_type_id = &info.variants[*tag];
let payload_ty = registry.get_type(payload_type_id)?;
let payload = value.to_ptr(arena, registry, payload_type_id)?;

// Undo the wrapper pointer added when the payload is memory
// allocated (e.g. a nested >=2-variant enum), so that the copy
// below reads the payload data rather than the wrapper pointer.
let payload = if payload_ty.is_memory_allocated(registry)? {
*payload.cast::<NonNull<()>>().as_ref()
} else {
payload
};

let (layout, tag_layout, variant_layouts) =
crate::types::r#enum::get_layout_for_variants(
registry,
Expand Down
29 changes: 29 additions & 0 deletions test_data/programs/nested_enum_arg.cairo
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
// Both `Inner` and `Outer` are 2-variant enums, so both are memory-allocated.
// Passing an `Outer` as an argument forces `to_ptr` to serialize a nested
// memory-allocated enum. The return value depends on both discriminants and the
// payload, so any corruption of the nested value changes the result.

#[derive(Drop)]
enum Inner {
A: felt252,
B: felt252,
}

#[derive(Drop)]
enum Outer {
X: Inner,
Y: Inner,
}

fn run_test(x: Outer) -> felt252 {
match x {
Outer::X(inner) => match inner {
Inner::A(v) => v,
Inner::B(v) => v + 1,
},
Outer::Y(inner) => match inner {
Inner::A(v) => v + 2,
Inner::B(v) => v + 3,
},
}
}
49 changes: 49 additions & 0 deletions tests/tests/enums.rs
Original file line number Diff line number Diff line change
Expand Up @@ -50,3 +50,52 @@ fn single_variant_enum_in_array_matches_vm() {
)
.expect("single-variant enum in array must agree between VM and native");
}

#[test]
fn nested_enum_argument_matches_vm() {
let program = &load_program_and_runner("programs/nested_enum_arg");

// (outer_tag, inner_tag). Both enums have 2 variants, so the tag equals the
// variant index (no selector encoding is needed).
for (outer_tag, inner_tag) in [(0, 0), (0, 1), (1, 0), (1, 1)] {
let payload = Felt::from(0x1234);

let result_vm = run_vm_program(
program,
"run_test",
vec![
Arg::Value(Felt::from(outer_tag as u64)),
Arg::Value(Felt::from(inner_tag as u64)),
Arg::Value(payload),
],
Some(DEFAULT_GAS as usize),
)
.unwrap();

let result_native = run_native_program(
program,
"run_test",
&[Value::Enum {
tag: outer_tag,
value: Box::new(Value::Enum {
tag: inner_tag,
value: Box::new(Value::Felt252(payload)),
debug_name: None,
}),
debug_name: None,
}],
Some(DEFAULT_GAS),
Option::<DummySyscallHandler>::None,
);

compare_outputs(
&program.1,
&program.2.find_function("run_test").unwrap().id,
&result_vm,
&result_native,
)
.unwrap_or_else(|e| {
panic!("nested enum (outer={outer_tag}, inner={inner_tag}) mismatch: {e:?}")
});
}
}
Loading