Skip to content

Commit b7092eb

Browse files
committed
Support enum in codegen
1 parent ee8e0e2 commit b7092eb

3 files changed

Lines changed: 78 additions & 6 deletions

File tree

crates/codegen/src/db/queries/abi.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -201,6 +201,7 @@ pub fn abi_type(db: &dyn CodegenDb, ty: TypeId) -> AbiType {
201201
| ir::TypeKind::Contract(_)
202202
| ir::TypeKind::Map(_)
203203
| ir::TypeKind::MPtr(_)
204+
| ir::TypeKind::Enum(_)
204205
| ir::TypeKind::SPtr(_) => unreachable!(),
205206
}
206207
}

crates/codegen/src/yul/runtime/data.rs

Lines changed: 54 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@ use crate::{
99

1010
use super::{DefaultRuntimeProvider, RuntimeFunction, RuntimeProvider};
1111

12-
use fe_mir::ir::TypeId;
12+
use fe_mir::ir::{types::TupleDef, Type, TypeId, TypeKind};
1313

1414
use yultsur::*;
1515

@@ -287,6 +287,59 @@ pub(super) fn make_aggregate_init(
287287
RuntimeFunction(func_def)
288288
}
289289

290+
pub(super) fn make_enum_init(
291+
provider: &mut DefaultRuntimeProvider,
292+
db: &dyn CodegenDb,
293+
func_name: &str,
294+
legalized_ty: TypeId,
295+
arg_tys: Vec<TypeId>,
296+
) -> RuntimeFunction {
297+
debug_assert!(arg_tys.len() > 1);
298+
299+
let func_name = YulVariable::new(func_name);
300+
let is_sptr = legalized_ty.is_sptr(db.upcast());
301+
let ptr = YulVariable::new("ptr");
302+
let tag = YulVariable::new("tag");
303+
let tag_ty = arg_tys[0];
304+
let enum_data = || {
305+
(0..arg_tys.len() - 1)
306+
.into_iter()
307+
.map(|i| YulVariable::new(format! {"arg{}", i}))
308+
};
309+
310+
let tuple_def = TupleDef {
311+
items: arg_tys.iter().copied().skip(1).collect(),
312+
};
313+
let tuple_ty = db.mir_intern_type(
314+
Type {
315+
kind: TypeKind::Tuple(tuple_def),
316+
analyzer_ty: None,
317+
}
318+
.into(),
319+
);
320+
let data_ptr_ty = make_ptr(db, tuple_ty, is_sptr);
321+
let data_offset = legalized_ty
322+
.deref(db.upcast())
323+
.enum_data_offset(db.upcast(), SLOT_SIZE);
324+
let enum_data_init = statements! {
325+
[statement! {[ptr.ident()] := add([ptr.expr()], [literal_expression!{(data_offset)}])}]
326+
[yul::Statement::Expression(provider.aggregate_init(
327+
db,
328+
ptr.expr(),
329+
enum_data().map(|arg| arg.expr()).collect(),
330+
data_ptr_ty, arg_tys.iter().copied().skip(1).collect()))]
331+
};
332+
333+
let enum_data_args: Vec<_> = enum_data().map(|var| var.ident()).collect();
334+
let func_def = function_definition! {
335+
function [func_name.ident()]([ptr.ident()], [tag.ident()], [enum_data_args...]) {
336+
[yul::Statement::Expression(provider.ptr_store(db, ptr.expr(), tag.expr(), make_ptr(db, tag_ty, is_sptr)))]
337+
[enum_data_init...]
338+
}
339+
};
340+
RuntimeFunction::from_statement(func_def)
341+
}
342+
290343
pub(super) fn make_string_copy(
291344
provider: &mut DefaultRuntimeProvider,
292345
db: &dyn CodegenDb,

crates/codegen/src/yul/runtime/mod.rs

Lines changed: 23 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -367,17 +367,35 @@ impl RuntimeProvider for DefaultRuntimeProvider {
367367
&mut self,
368368
db: &dyn CodegenDb,
369369
ptr: yul::Expression,
370-
args: Vec<yul::Expression>,
370+
mut args: Vec<yul::Expression>,
371371
ptr_ty: TypeId,
372372
arg_tys: Vec<TypeId>,
373373
) -> yul::Expression {
374374
debug_assert!(ptr_ty.is_ptr(db.upcast()));
375-
let name = format!("$aggregate_init_{}", ptr_ty.0);
375+
let deref_ty = ptr_ty.deref(db.upcast());
376+
377+
// Handle unit enum variant.
378+
if args.len() == 1 && deref_ty.is_enum(db.upcast()) {
379+
let tag = args.pop().unwrap();
380+
let tag_ty = arg_tys[0];
381+
let is_sptr = ptr_ty.is_sptr(db.upcast());
382+
return self.ptr_store(db, ptr, tag, make_ptr(db, tag_ty, is_sptr));
383+
}
384+
385+
let deref_ty = ptr_ty.deref(db.upcast());
376386
let args = std::iter::once(ptr).chain(args.into_iter()).collect();
377387
let legalized_ty = db.codegen_legalized_type(ptr_ty);
378-
self.create_then_call(&name, args, |provider| {
379-
data::make_aggregate_init(provider, db, &name, legalized_ty, arg_tys)
380-
})
388+
if deref_ty.is_enum(db.upcast()) {
389+
let name = format!("enum_init_{}_{}", ptr_ty.0, arg_tys[1].0);
390+
self.create_then_call(&name, args, |provider| {
391+
data::make_enum_init(provider, db, &name, legalized_ty, arg_tys)
392+
})
393+
} else {
394+
let name = format!("$aggregate_init_{}", ptr_ty.0);
395+
self.create_then_call(&name, args, |provider| {
396+
data::make_aggregate_init(provider, db, &name, legalized_ty, arg_tys)
397+
})
398+
}
381399
}
382400

383401
fn string_copy(

0 commit comments

Comments
 (0)