Skip to content

Commit 7da5c73

Browse files
authored
Support variants with multiple fields in UDTs (#850)
1 parent bb98641 commit 7da5c73

2 files changed

Lines changed: 250 additions & 80 deletions

File tree

soroban-sdk-macros/src/derive_enum.rs

Lines changed: 218 additions & 79 deletions
Original file line numberDiff line numberDiff line change
@@ -5,8 +5,8 @@ use soroban_env_common::Symbol;
55
use syn::{spanned::Spanned, Attribute, DataEnum, Error, Fields, Ident, Path};
66

77
use stellar_xdr::{
8-
ScSpecEntry, ScSpecTypeDef, ScSpecUdtUnionCaseTupleV0, ScSpecUdtUnionCaseV0,
9-
ScSpecUdtUnionCaseVoidV0, ScSpecUdtUnionV0, StringM, WriteXdr,
8+
Error as XdrError, ScSpecEntry, ScSpecTypeDef, ScSpecUdtUnionCaseTupleV0, ScSpecUdtUnionCaseV0,
9+
ScSpecUdtUnionCaseVoidV0, ScSpecUdtUnionV0, StringM, VecM, WriteXdr,
1010
};
1111

1212
use crate::{doc::docs_from_attrs, map_type::map_type};
@@ -33,28 +33,19 @@ pub fn derive_type_enum(
3333
.iter()
3434
.map(|v| {
3535
// TODO: Choose discriminant type based on repr type of enum.
36-
// TODO: Should we use variants explicit discriminant? Probably not.
37-
// Should have a separate derive for those types of enums that maps
38-
// to an integer type only.
3936
// TODO: Use attributes tagged on variant to control whether field is included.
40-
// TODO: Support multi-field enum variants.
41-
// TODO: Or, error on multi-field enum variants.
4237
// TODO: Handle field names longer than a symbol. Hash the name? Truncate the name?
4338
let ident = &v.ident;
4439
let name = &ident.to_string();
4540
if let Err(e) = Symbol::try_from_str(name) {
4641
errors.push(Error::new(ident.span(), format!("enum variant name {}", e)));
4742
}
48-
if v.fields.len() > 1 {
49-
errors.push(Error::new(v.fields.span(), format!("enum variant name {} has too many tuple values, max 1 supported", ident)));
50-
}
5143
match v.fields {
5244
Fields::Named(_) => {
5345
errors.push(Error::new(v.fields.span(), format!("enum variant {} has unsupported named fields", ident)));
5446
}
55-
_ => {}
56-
};
57-
let field = v.fields.iter().next();
47+
_ => { }
48+
}
5849
let discriminant_const_sym_ident = format_ident!("DISCRIMINANT_SYM_{}", name.to_uppercase());
5950
let discriminant_const_u64_ident = format_ident!("DISCRIMINANT_U64_{}", name.to_uppercase());
6051
let discriminant_const_sym = quote! {
@@ -67,75 +58,34 @@ pub fn derive_type_enum(
6758
#discriminant_const_sym
6859
#discriminant_const_u64
6960
};
70-
if let Some(f) = field {
71-
let spec_case = ScSpecUdtUnionCaseV0::TupleV0(
72-
ScSpecUdtUnionCaseTupleV0 {
73-
doc: docs_from_attrs(&v.attrs).try_into().unwrap(), // TODO: Truncate docs, or display friendly compile error.
74-
name: name.try_into().unwrap_or_else(|_| StringM::default()),
75-
type_: vec![
76-
match map_type(&f.ty) {
77-
Ok(t) => t,
78-
Err(e) => {
79-
errors.push(e);
80-
ScSpecTypeDef::I32
81-
}
82-
},
83-
].try_into().unwrap()
84-
}
61+
let has_fields = v.fields.iter().next().is_some();
62+
if has_fields {
63+
let VariantTokens {
64+
spec_case, try_from, try_into, try_from_xdr, into_xdr
65+
} = map_tuple_variant(
66+
path,
67+
enum_ident,
68+
&name,
69+
ident,
70+
&v.attrs,
71+
&discriminant_const_sym_ident,
72+
&discriminant_const_u64_ident,
73+
&v.fields,
74+
&mut errors,
8575
);
86-
let try_from = quote! {
87-
#discriminant_const_u64_ident => {
88-
if iter.len() > 1 {
89-
return Err(#path::ConversionError);
90-
}
91-
Self::#ident(iter.next().ok_or(#path::ConversionError)??.try_into_val(env)?)
92-
}
93-
};
94-
let try_into = quote! {
95-
#enum_ident::#ident(ref value) => {
96-
let tup: (#path::RawVal, #path::RawVal) = (#discriminant_const_sym_ident.into(), value.try_into_val(env)?);
97-
tup.try_into_val(env)
98-
}
99-
};
100-
let try_from_xdr = quote! {
101-
#name => {
102-
if iter.len() > 1 {
103-
return Err(#path::xdr::Error::Invalid);
104-
}
105-
let rv: #path::RawVal = iter.next().ok_or(#path::xdr::Error::Invalid)?.try_into_val(env).map_err(|_| #path::xdr::Error::Invalid)?;
106-
Self::#ident(rv.try_into_val(env).map_err(|_| #path::xdr::Error::Invalid)?)
107-
}
108-
};
109-
let into_xdr = quote! { #enum_ident::#ident(value) => (#name, value).try_into().map_err(|_| #path::xdr::Error::Invalid)? };
11076
(spec_case, discriminant_const, try_from, try_into, try_from_xdr, into_xdr)
11177
} else {
112-
let spec_case = ScSpecUdtUnionCaseV0::VoidV0(ScSpecUdtUnionCaseVoidV0 {
113-
doc: docs_from_attrs(&v.attrs).try_into().unwrap(), // TODO: Truncate docs, or display friendly compile error.
114-
name: name.try_into().unwrap_or_else(|_| StringM::default()),
115-
});
116-
let try_from = quote! {
117-
#discriminant_const_u64_ident => {
118-
if iter.len() > 0 {
119-
return Err(#path::ConversionError);
120-
}
121-
Self::#ident
122-
}
123-
};
124-
let try_into = quote! {
125-
#enum_ident::#ident => {
126-
let tup: (#path::RawVal,) = (#discriminant_const_sym_ident.into(),);
127-
tup.try_into_val(env)
128-
}
129-
};
130-
let try_from_xdr = quote! {
131-
#name => {
132-
if iter.len() > 0 {
133-
return Err(#path::xdr::Error::Invalid);
134-
}
135-
Self::#ident
136-
}
137-
};
138-
let into_xdr = quote! { #enum_ident::#ident => (#name,).try_into().map_err(|_| #path::xdr::Error::Invalid)? };
78+
let VariantTokens {
79+
spec_case, try_from, try_into, try_from_xdr, into_xdr
80+
} = map_empty_variant(
81+
path,
82+
enum_ident,
83+
&name,
84+
ident,
85+
&v.attrs,
86+
&discriminant_const_sym_ident,
87+
&discriminant_const_u64_ident,
88+
);
13989
(spec_case, discriminant_const, try_from, try_into, try_from_xdr, into_xdr)
14090
}
14191
})
@@ -309,3 +259,192 @@ pub fn derive_type_enum(
309259
}
310260
}
311261
}
262+
263+
struct VariantTokens {
264+
spec_case: ScSpecUdtUnionCaseV0,
265+
try_from: TokenStream2,
266+
try_into: TokenStream2,
267+
try_from_xdr: TokenStream2,
268+
into_xdr: TokenStream2,
269+
}
270+
271+
fn map_empty_variant(
272+
path: &Path,
273+
enum_ident: &Ident,
274+
name: &str,
275+
ident: &Ident,
276+
attrs: &[Attribute],
277+
discriminant_const_sym_ident: &Ident,
278+
discriminant_const_u64_ident: &Ident,
279+
) -> VariantTokens {
280+
let spec_case = ScSpecUdtUnionCaseV0::VoidV0(ScSpecUdtUnionCaseVoidV0 {
281+
doc: docs_from_attrs(attrs).try_into().unwrap(), // TODO: Truncate docs, or display friendly compile error.
282+
name: name.try_into().unwrap_or_else(|_| StringM::default()),
283+
});
284+
let try_from = quote! {
285+
#discriminant_const_u64_ident => {
286+
if iter.len() > 0 {
287+
return Err(#path::ConversionError);
288+
}
289+
Self::#ident
290+
}
291+
};
292+
let try_into = quote! {
293+
#enum_ident::#ident => {
294+
let tup: (#path::RawVal,) = (#discriminant_const_sym_ident.into(),);
295+
tup.try_into_val(env)
296+
}
297+
};
298+
let try_from_xdr = quote! {
299+
#name => {
300+
if iter.len() > 0 {
301+
return Err(#path::xdr::Error::Invalid);
302+
}
303+
Self::#ident
304+
}
305+
};
306+
let into_xdr = quote! { #enum_ident::#ident => (#name,).try_into().map_err(|_| #path::xdr::Error::Invalid)? };
307+
308+
VariantTokens {
309+
spec_case,
310+
try_from,
311+
try_into,
312+
try_from_xdr,
313+
into_xdr,
314+
}
315+
}
316+
317+
fn map_tuple_variant(
318+
path: &Path,
319+
enum_ident: &Ident,
320+
name: &str,
321+
ident: &Ident,
322+
attrs: &[Attribute],
323+
discriminant_const_sym_ident: &Ident,
324+
discriminant_const_u64_ident: &Ident,
325+
fields: &Fields,
326+
errors: &mut Vec<Error>,
327+
) -> VariantTokens {
328+
let spec_case = {
329+
let field_types = fields
330+
.iter()
331+
.map(|f| match map_type(&f.ty) {
332+
Ok(t) => t,
333+
Err(e) => {
334+
errors.push(e);
335+
ScSpecTypeDef::I32
336+
}
337+
})
338+
.collect::<Vec<_>>();
339+
let field_types = match VecM::try_from(field_types) {
340+
Ok(t) => t,
341+
Err(e) => {
342+
let v = VecM::default();
343+
let max_len = v.max_len();
344+
match e {
345+
XdrError::LengthExceedsMax => {
346+
errors.push(Error::new(
347+
fields.span(),
348+
format!(
349+
"enum variant name {} has too many tuple values, max {} supported",
350+
ident, max_len
351+
),
352+
));
353+
}
354+
e => {
355+
errors.push(Error::new(fields.span(), format!("{e}")));
356+
}
357+
}
358+
v
359+
}
360+
};
361+
ScSpecUdtUnionCaseV0::TupleV0(ScSpecUdtUnionCaseTupleV0 {
362+
doc: docs_from_attrs(attrs).try_into().unwrap(), // TODO: Truncate docs, or display friendly compile error.
363+
name: name.try_into().unwrap_or_else(|_| StringM::default()),
364+
type_: field_types.try_into().unwrap(),
365+
})
366+
};
367+
let num_fields = fields.iter().len();
368+
let try_from = {
369+
let field_convs = fields
370+
.iter()
371+
.enumerate()
372+
.map(|(_i, _f)| {
373+
quote! {
374+
iter.next().ok_or(#path::ConversionError)??.try_into_val(env)?
375+
}
376+
})
377+
.collect::<Vec<_>>();
378+
quote! {
379+
#discriminant_const_u64_ident => {
380+
if iter.len() > #num_fields {
381+
return Err(#path::ConversionError);
382+
}
383+
Self::#ident( #(#field_convs,)* )
384+
}
385+
}
386+
};
387+
let try_into = {
388+
let fragments = fields
389+
.iter()
390+
.enumerate()
391+
.map(|(i, _f)| {
392+
let binding_name = format_ident!("value{i}");
393+
let field_conv = quote! {
394+
#binding_name.try_into_val(env)?
395+
};
396+
let tup_elem_type = quote! {
397+
#path::RawVal
398+
};
399+
(binding_name, field_conv, tup_elem_type)
400+
})
401+
.multiunzip();
402+
let (binding_names, field_convs, tup_elem_types): (Vec<_>, Vec<_>, Vec<_>) = fragments;
403+
quote! {
404+
#enum_ident::#ident(#(ref #binding_names,)* ) => {
405+
let tup: (#path::RawVal, #(#tup_elem_types,)* ) = (#discriminant_const_sym_ident.into(), #(#field_convs,)* );
406+
tup.try_into_val(env)
407+
}
408+
}
409+
};
410+
let try_from_xdr = {
411+
let fragments = fields.iter().enumerate().map(|(i, _f)| {
412+
let rawval_name = format_ident!("rv{i}");
413+
let rawval_binding = quote! {
414+
let #rawval_name: #path::RawVal = iter.next().ok_or(#path::xdr::Error::Invalid)?.try_into_val(env).map_err(|_| #path::xdr::Error::Invalid)?;
415+
};
416+
let into_field = quote! {
417+
#rawval_name.try_into_val(env).map_err(|_| #path::xdr::Error::Invalid)?
418+
};
419+
(rawval_binding, into_field)
420+
}).multiunzip();
421+
let (rawval_bindings, into_fields): (Vec<_>, Vec<_>) = fragments;
422+
quote! {
423+
#name => {
424+
if iter.len() > #num_fields {
425+
return Err(#path::xdr::Error::Invalid);
426+
}
427+
#(#rawval_bindings)*
428+
Self::#ident( #(#into_fields,)* )
429+
}
430+
}
431+
};
432+
let into_xdr = {
433+
let binding_names = fields
434+
.iter()
435+
.enumerate()
436+
.map(|(i, _f)| format_ident!("value{i}"))
437+
.collect::<Vec<_>>();
438+
quote! {
439+
#enum_ident::#ident( #(#binding_names,)* ) => (#name, #(#binding_names,)* ).try_into().map_err(|_| #path::xdr::Error::Invalid)?
440+
}
441+
};
442+
443+
VariantTokens {
444+
spec_case,
445+
try_from,
446+
try_into,
447+
try_from_xdr,
448+
into_xdr,
449+
}
450+
}

soroban-sdk/src/tests/contract_udt_enum.rs

Lines changed: 32 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,23 @@
11
use crate as soroban_sdk;
2+
use soroban_sdk::xdr::ScVec;
23
use soroban_sdk::{
3-
contractimpl, contracttype, symbol, vec, ConversionError, Env, IntoVal, RawVal, TryFromVal, Vec,
4+
contractimpl, contracttype, symbol, vec, ConversionError, Env, IntoVal, RawVal, TryFromVal,
5+
TryIntoVal, Vec,
46
};
57

68
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
79
#[contracttype]
810
pub enum Udt {
911
Aaa,
1012
Bbb(i32),
13+
MaxFields(u32, u32, u32, u32, u32, u32, u32, u32, u32, u32, u32, u32),
14+
Nested(Udt2, Udt2),
15+
}
16+
17+
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
18+
#[contracttype]
19+
pub struct Udt2 {
20+
a: u32,
1121
}
1222

1323
pub struct Contract;
@@ -56,3 +66,24 @@ fn test_error_on_partial_decode() {
5666
let udt = Udt::try_from_val(&env, &vec.to_raw());
5767
assert_eq!(udt, Err(ConversionError));
5868
}
69+
70+
#[test]
71+
fn round_trips() {
72+
let env = Env::default();
73+
74+
let before = Udt::Nested(Udt2 { a: 1 }, Udt2 { a: 2 });
75+
let rawval: RawVal = before.try_into_val(&env).unwrap();
76+
let after: Udt = rawval.try_into_val(&env).unwrap();
77+
assert_eq!(before, after);
78+
let scvec: ScVec = before.try_into().unwrap();
79+
let after: Udt = scvec.try_into_val(&env).unwrap();
80+
assert_eq!(before, after);
81+
82+
let before = Udt::MaxFields(1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12);
83+
let rawval: RawVal = before.try_into_val(&env).unwrap();
84+
let after: Udt = rawval.try_into_val(&env).unwrap();
85+
assert_eq!(before, after);
86+
let scvec: ScVec = before.try_into().unwrap();
87+
let after: Udt = scvec.try_into_val(&env).unwrap();
88+
assert_eq!(before, after);
89+
}

0 commit comments

Comments
 (0)