Skip to main content

autophagy_mlir/
compiler.rs

1use crate::Error;
2use autophagy::{Fn, Struct};
3use melior::{
4    dialect::{arith, func, llvm, memref, scf},
5    ir::{
6        attribute::{
7            DenseI64ArrayAttribute, FlatSymbolRefAttribute, FloatAttribute, IntegerAttribute,
8            StringAttribute, TypeAttribute,
9        },
10        r#type::{FunctionType, IntegerType, MemRefType},
11        Attribute, Block, Identifier, Location, Module, OperationRef, Region, ShapedTypeLike, Type,
12        TypeLike, Value, ValueLike,
13    },
14    Context,
15};
16use std::collections::HashMap;
17use train_map::TrainMap;
18
19struct StructInfo<'c> {
20    r#type: Type<'c>,
21    field_types: Vec<Type<'c>>,
22    field_indices: HashMap<String, usize>,
23}
24
25pub struct Compiler<'c, 'm> {
26    context: &'c Context,
27    module: &'m Module<'c>,
28    functions: HashMap<String, FunctionType<'c>>,
29    structs: HashMap<String, StructInfo<'c>>,
30}
31
32impl<'c, 'm> Compiler<'c, 'm> {
33    pub fn new(context: &'c Context, module: &'m Module<'c>) -> Self {
34        Self {
35            context,
36            module,
37            functions: Default::default(),
38            structs: Default::default(),
39        }
40    }
41
42    pub fn compile_struct(&mut self, r#struct: &Struct) -> Result<(), Error> {
43        let types = r#struct
44            .ast()
45            .fields
46            .iter()
47            .map(|field| self.compile_type(&field.ty))
48            .collect::<Result<Vec<_>, _>>()?;
49
50        self.structs.insert(
51            r#struct.name().into(),
52            StructInfo {
53                r#type: llvm::r#type::r#struct(self.context, &types, false),
54                field_types: types,
55                field_indices: r#struct
56                    .ast()
57                    .fields
58                    .iter()
59                    .enumerate()
60                    .flat_map(|(index, field)| {
61                        field.ident.as_ref().map(|ident| (ident.to_string(), index))
62                    })
63                    .collect(),
64            },
65        );
66
67        Ok(())
68    }
69
70    pub fn compile_fn(&mut self, r#fn: &Fn) -> Result<(), Error> {
71        let function = r#fn.ast();
72        let context = self.context;
73        let location = Location::unknown(context);
74        let argument_types = function
75            .sig
76            .inputs
77            .iter()
78            .map(|argument| match argument {
79                syn::FnArg::Typed(typed) => self.compile_type(&typed.ty),
80                syn::FnArg::Receiver(_) => Err(Error::NotSupported("self receiver")),
81            })
82            .collect::<Result<Vec<_>, _>>()?;
83        let result_types = match &function.sig.output {
84            syn::ReturnType::Default => vec![],
85            syn::ReturnType::Type(_, r#type) => vec![self.compile_type(r#type)?],
86        };
87        let mut variables = TrainMap::new();
88
89        let name = function.sig.ident.to_string();
90        let function_type = FunctionType::new(context, &argument_types, &result_types);
91        self.functions.insert(name.clone(), function_type);
92
93        self.module.body().append_operation(func::func(
94            context,
95            StringAttribute::new(context, &name),
96            TypeAttribute::new(function_type.into()),
97            {
98                let block = Block::new(
99                    &argument_types
100                        .iter()
101                        .map(|&r#type| (r#type, location))
102                        .collect::<Vec<_>>(),
103                );
104
105                for result in function.sig.inputs.iter().map(|argument| match argument {
106                    syn::FnArg::Typed(typed) => match typed.pat.as_ref() {
107                        syn::Pat::Ident(identifier) => {
108                            Ok((identifier.ident.to_string(), &typed.ty))
109                        }
110                        _ => Err(Error::NotSupported("non-identifier pattern")),
111                    },
112                    syn::FnArg::Receiver(_) => Err(Error::NotSupported("self receiver")),
113                }) {
114                    let (name, r#type) = result?;
115
116                    let ptr = block
117                        .append_operation(memref::alloca(
118                            context,
119                            MemRefType::new(self.compile_type(r#type)?, &[], None, None),
120                            &[],
121                            &[],
122                            None,
123                            location,
124                        ))
125                        .result(0)?
126                        .into();
127
128                    block.append_operation(memref::store(
129                        block.argument(0)?.into(),
130                        ptr,
131                        &[],
132                        location,
133                    ));
134
135                    variables.insert(name, ptr);
136                }
137
138                self.compile_statements(&block, &function.block.stmts, true, &mut variables)?;
139
140                let region = Region::new();
141                region.append_block(block);
142                region
143            },
144            &[(
145                Identifier::new(context, "llvm.emit_c_interface"),
146                Attribute::unit(context),
147            )],
148            location,
149        ));
150
151        Ok(())
152    }
153
154    fn compile_type(&self, r#type: &syn::Type) -> Result<Type<'c>, Error> {
155        Ok(match r#type {
156            syn::Type::Path(path) => {
157                if let Some(identifier) = path.path.get_ident() {
158                    self.compile_primitive_type(&identifier.to_string())?
159                } else {
160                    return Err(Error::NotSupported("custom type"));
161                }
162            }
163            syn::Type::Reference(reference) => {
164                MemRefType::new(self.compile_type(&reference.elem)?, &[], None, None).into()
165            }
166            _ => todo!(),
167        })
168    }
169
170    fn compile_primitive_type(&self, name: &str) -> Result<Type<'c>, Error> {
171        let context = self.context;
172
173        Ok(match name {
174            "bool" => IntegerType::new(context, 1).into(),
175            "f32" => Type::float32(context),
176            "f64" => Type::float64(context),
177            "isize" | "usize" => Type::index(context),
178            "i8" | "u8" => IntegerType::new(context, 8).into(),
179            "i16" | "u16" => IntegerType::new(context, 16).into(),
180            "i32" | "u32" => IntegerType::new(context, 32).into(),
181            "i64" | "u64" => IntegerType::new(context, 64).into(),
182            name => {
183                self.structs
184                    .get(name)
185                    .ok_or_else(|| Error::TypeNotDefined(name.into()))?
186                    .r#type
187            }
188        })
189    }
190
191    fn compile_block(
192        &self,
193        block: &syn::Block,
194        function_scope: bool,
195        variables: &mut TrainMap<String, Value<'c, '_>>,
196    ) -> Result<Region<'c>, Error> {
197        Ok(self
198            .compile_block_expression(block, function_scope, variables)?
199            .0)
200    }
201
202    fn compile_block_expression(
203        &self,
204        block: &syn::Block,
205        function_scope: bool,
206        variables: &mut TrainMap<String, Value<'c, '_>>,
207    ) -> Result<(Region<'c>, Option<Type<'c>>), Error> {
208        let builder = Block::new(&[]);
209        let mut variables = variables.fork();
210
211        let r#type =
212            self.compile_statements(&builder, &block.stmts, function_scope, &mut variables)?;
213
214        let region = Region::new();
215        region.append_block(builder);
216        Ok((region, r#type))
217    }
218
219    fn compile_statements<'a>(
220        &self,
221        builder: &'a Block<'c>,
222        statements: &[syn::Stmt],
223        function_scope: bool,
224        variables: &mut TrainMap<String, Value<'c, 'a>>,
225    ) -> Result<Option<Type<'c>>, Error> {
226        let context = self.context;
227        let location = Location::unknown(context);
228        let terminator = if function_scope {
229            func::r#return
230        } else {
231            scf::r#yield
232        };
233        let mut return_value = None;
234
235        for statement in statements {
236            match statement {
237                syn::Stmt::Local(local) => self.compile_local_binding(builder, local, variables)?,
238                syn::Stmt::Item(_) => return Err(Error::NotSupported("local item definition")),
239                syn::Stmt::Expr(expression, semicolon) => {
240                    let value = self.compile_expression(builder, expression, variables)?;
241
242                    if semicolon.is_none() {
243                        return_value = value;
244                    }
245                }
246                syn::Stmt::Macro(_) => return Err(Error::NotSupported("macro")),
247            }
248        }
249
250        builder.append_operation(if let Some(value) = return_value {
251            terminator(&[value], location)
252        } else {
253            terminator(&[], location)
254        });
255
256        Ok(if function_scope {
257            None
258        } else {
259            return_value.map(|value| value.r#type())
260        })
261    }
262
263    fn compile_local_binding<'a>(
264        &self,
265        builder: &'a Block<'c>,
266        local: &syn::Local,
267        variables: &mut TrainMap<String, Value<'c, 'a>>,
268    ) -> Result<(), Error> {
269        let context = self.context;
270
271        let value = self.compile_expression_value(
272            builder,
273            if let Some(initial) = &local.init {
274                &initial.expr
275            } else {
276                return Err(Error::NotSupported("uninitialized let binding"));
277            },
278            variables,
279        )?;
280        let ptr = builder
281            .append_operation(memref::alloca(
282                context,
283                MemRefType::new(value.r#type(), &[], None, None),
284                &[],
285                &[],
286                None,
287                Location::unknown(context),
288            ))
289            .result(0)?
290            .into();
291
292        builder.append_operation(memref::store(
293            value,
294            ptr,
295            &[],
296            Location::unknown(self.context),
297        ));
298
299        variables.insert(
300            match &local.pat {
301                syn::Pat::Ident(identifier) => identifier.ident.to_string(),
302                _ => return Err(Error::NotSupported("non-identifier pattern")),
303            },
304            ptr,
305        );
306
307        Ok(())
308    }
309
310    fn compile_expression_value<'a>(
311        &self,
312        builder: &'a Block<'c>,
313        expression: &syn::Expr,
314        variables: &mut TrainMap<String, Value<'c, 'a>>,
315    ) -> Result<Value<'c, 'a>, Error> {
316        Ok(
317            if let Some(value) = self.compile_expression(builder, expression, variables)? {
318                value
319            } else {
320                self.compile_unit(builder)?
321            },
322        )
323    }
324
325    fn compile_expression<'a>(
326        &self,
327        builder: &'a Block<'c>,
328        expression: &syn::Expr,
329        variables: &mut TrainMap<String, Value<'c, 'a>>,
330    ) -> Result<Option<Value<'c, 'a>>, Error> {
331        let context = self.context;
332        let location = Location::unknown(context);
333
334        Ok(match expression {
335            syn::Expr::Assign(assign) => {
336                // TODO Support a `*` pointer dereference.
337                // TODO Support recursive LHS dereference.
338                builder.append_operation(match assign.left.as_ref() {
339                    syn::Expr::Field(field) => {
340                        let mut ptr = self.compile_ptr(builder, &field.base, variables)?;
341
342                        while MemRefType::try_from(ptr.r#type())?.element().is_mem_ref() {
343                            ptr = builder
344                                .append_operation(memref::load(ptr, &[], location))
345                                .result(0)?
346                                .into();
347                        }
348
349                        let info =
350                            self.get_struct_info(MemRefType::try_from(ptr.r#type())?.element())?;
351
352                        memref::store(
353                            builder
354                                .append_operation(llvm::insert_value(
355                                    context,
356                                    builder
357                                        .append_operation(memref::load(ptr, &[], location))
358                                        .result(0)?
359                                        .into(),
360                                    DenseI64ArrayAttribute::new(
361                                        context,
362                                        &[self.get_struct_field_index(&field.member, info)? as i64],
363                                    ),
364                                    self.compile_expression_value(
365                                        builder,
366                                        &assign.right,
367                                        variables,
368                                    )?,
369                                    location,
370                                ))
371                                .result(0)?
372                                .into(),
373                            ptr,
374                            &[],
375                            location,
376                        )
377                    }
378                    syn::Expr::Path(path) => {
379                        let ptr = self.compile_variable(
380                            &self.convert_path_to_identifier(&path.path)?,
381                            variables,
382                        )?;
383
384                        memref::store(
385                            self.compile_expression_value(builder, &assign.right, variables)?,
386                            ptr,
387                            &[],
388                            location,
389                        )
390                    }
391                    _ => todo!(),
392                });
393
394                None
395            }
396            syn::Expr::Binary(operation) => Some(
397                self.compile_binary_operation(builder, operation, variables)?
398                    .result(0)?
399                    .into(),
400            ),
401            syn::Expr::Block(block) => {
402                let (region, r#type) =
403                    self.compile_block_expression(&block.block, false, variables)?;
404
405                Some(
406                    builder
407                        .append_operation(scf::execute_region(
408                            &r#type.into_iter().collect::<Vec<_>>(),
409                            region,
410                            location,
411                        ))
412                        .result(0)?
413                        .into(),
414                )
415            }
416            syn::Expr::Call(call) => {
417                let function = self.compile_expression_value(builder, &call.func, variables)?;
418
419                builder
420                    .append_operation(func::call_indirect(
421                        function,
422                        &call
423                            .args
424                            .iter()
425                            .map(|argument| {
426                                self.compile_expression_value(builder, argument, variables)
427                            })
428                            .collect::<Result<Vec<_>, _>>()?,
429                        &FunctionType::try_from(function.r#type())?
430                            .result(0)
431                            .into_iter()
432                            .collect::<Vec<_>>(),
433                        location,
434                    ))
435                    .result(0)
436                    .map(Into::into)
437                    .ok()
438            }
439            syn::Expr::Field(field) => {
440                let mut value = self
441                    .compile_expression(builder, &field.base, variables)?
442                    .ok_or_else(|| {
443                        Error::ValueExpected("struct field access requires struct value".into())
444                    })?;
445
446                while value.r#type().is_mem_ref() {
447                    value = builder
448                        .append_operation(memref::load(value, &[], location))
449                        .result(0)?
450                        .into();
451                }
452
453                let info = self.get_struct_info(value.r#type())?;
454                let index = self.get_struct_field_index(&field.member, info)?;
455
456                Some(
457                    builder
458                        .append_operation(llvm::extract_value(
459                            context,
460                            value,
461                            DenseI64ArrayAttribute::new(context, &[index as i64]),
462                            info.field_types[index],
463                            location,
464                        ))
465                        .result(0)?
466                        .into(),
467                )
468            }
469            syn::Expr::If(r#if) => {
470                let condition = self.compile_expression_value(builder, &r#if.cond, variables)?;
471                let (then_region, then_type) =
472                    self.compile_block_expression(&r#if.then_branch, false, variables)?;
473                let (else_region, else_type) = if let Some((_, expression)) = &r#if.else_branch {
474                    let block = Block::new(&[]);
475                    let mut variables = variables.fork();
476
477                    let value = self.compile_expression(&block, expression, &mut variables)?;
478                    block.append_operation(scf::r#yield(
479                        &value.into_iter().collect::<Vec<_>>(),
480                        location,
481                    ));
482
483                    let r#type = value.map(|value| value.r#type());
484                    let region = Region::new();
485                    region.append_block(block);
486
487                    (region, r#type)
488                } else {
489                    (Region::new(), None)
490                };
491
492                builder
493                    .append_operation(scf::r#if(
494                        condition,
495                        &then_type.or(else_type).into_iter().collect::<Vec<_>>(),
496                        then_region,
497                        else_region,
498                        location,
499                    ))
500                    .result(0)
501                    .map(Into::into)
502                    .ok()
503            }
504            syn::Expr::Lit(literal) => self
505                .compile_expression_literal(builder, literal)?
506                .result(0)
507                .map(Into::into)
508                .ok(),
509            syn::Expr::Loop(r#loop) => {
510                builder.append_operation(scf::r#while(
511                    &[],
512                    &[],
513                    {
514                        let block = Block::new(&[]);
515
516                        block.append_operation(scf::condition(
517                            block
518                                .append_operation(arith::constant(
519                                    context,
520                                    IntegerAttribute::new(
521                                        IntegerType::new(context, 1).into(),
522                                        true as i64,
523                                    )
524                                    .into(),
525                                    location,
526                                ))
527                                .result(0)?
528                                .into(),
529                            &[],
530                            location,
531                        ));
532
533                        let region = Region::new();
534                        region.append_block(block);
535                        region
536                    },
537                    self.compile_block(&r#loop.body, false, variables)?,
538                    location,
539                ));
540
541                None
542            }
543            syn::Expr::Paren(parenthesis) => {
544                self.compile_expression(builder, &parenthesis.expr, variables)?
545            }
546            syn::Expr::Path(path) => Some(self.compile_path(builder, path, variables)?),
547            syn::Expr::Struct(r#struct) => {
548                let name = self.convert_path_to_identifier(&r#struct.path)?;
549                let info = self
550                    .structs
551                    .get(&name)
552                    .ok_or(Error::StructNotDefined(name))?;
553                let mut value = builder
554                    .append_operation(llvm::undef(info.r#type, location))
555                    .result(0)?
556                    .into();
557
558                for field in &r#struct.fields {
559                    let index = self.get_struct_field_index(&field.member, info)?;
560
561                    value = builder
562                        .append_operation(llvm::insert_value(
563                            context,
564                            value,
565                            DenseI64ArrayAttribute::new(context, &[index as i64]),
566                            self.compile_expression_value(builder, &field.expr, variables)?,
567                            location,
568                        ))
569                        .result(0)?
570                        .into();
571                }
572
573                Some(value)
574            }
575            syn::Expr::Unary(operation) => self
576                .compile_unary_operation(builder, operation, variables)?
577                .result(0)
578                .map(Into::into)
579                .ok(),
580            syn::Expr::While(r#while) => {
581                builder.append_operation(scf::r#while(
582                    &[],
583                    &[],
584                    {
585                        let block = Block::new(&[]);
586                        let mut variables = variables.fork();
587
588                        block.append_operation(scf::condition(
589                            self.compile_expression_value(&block, &r#while.cond, &mut variables)?,
590                            &[],
591                            location,
592                        ));
593
594                        let region = Region::new();
595                        region.append_block(block);
596                        region
597                    },
598                    self.compile_block(&r#while.body, false, variables)?,
599                    location,
600                ));
601
602                None
603            }
604            _ => todo!(),
605        })
606    }
607
608    fn compile_ptr<'a>(
609        &self,
610        builder: &'a Block<'c>,
611        expression: &syn::Expr,
612        variables: &mut TrainMap<String, Value<'c, 'a>>,
613    ) -> Result<Value<'c, 'a>, Error> {
614        Ok(match expression {
615            syn::Expr::Path(path) => {
616                self.compile_variable(&self.convert_path_to_identifier(&path.path)?, variables)?
617            }
618            _ => self.compile_expression_value(builder, expression, variables)?,
619        })
620    }
621
622    fn compile_unary_operation<'a>(
623        &self,
624        builder: &'a Block<'c>,
625        operation: &syn::ExprUnary,
626        variables: &mut TrainMap<String, Value<'c, 'a>>,
627    ) -> Result<OperationRef<'c, 'a>, Error> {
628        let context = self.context;
629        let location = Location::unknown(context);
630        let value = self.compile_expression_value(builder, &operation.expr, variables)?;
631
632        // spell-checker: disable
633        Ok(builder.append_operation(match &operation.op {
634            syn::UnOp::Deref(_) => memref::load(value, &[], location),
635            syn::UnOp::Neg(_) => arith::subi(
636                builder
637                    .append_operation(arith::constant(
638                        context,
639                        IntegerAttribute::new(Type::index(context), 0).into(),
640                        location,
641                    ))
642                    .result(0)?
643                    .into(),
644                value,
645                location,
646            ),
647            syn::UnOp::Not(_) => arith::xori(
648                builder
649                    .append_operation(arith::constant(
650                        context,
651                        IntegerAttribute::new(Type::index(context), 0).into(),
652                        location,
653                    ))
654                    .result(0)?
655                    .into(),
656                value,
657                location,
658            ),
659            _ => return Err(Error::NotSupported("unknown unary operator")),
660        }))
661        // spell-checker: enable
662    }
663
664    fn compile_binary_operation<'a>(
665        &self,
666        builder: &'a Block<'c>,
667        operation: &syn::ExprBinary,
668        variables: &mut TrainMap<String, Value<'c, 'a>>,
669    ) -> Result<OperationRef<'c, 'a>, Error> {
670        let context = self.context;
671        let location = Location::unknown(context);
672        let left = self.compile_expression_value(builder, &operation.left, variables)?;
673        let right = self.compile_expression_value(builder, &operation.right, variables)?;
674
675        // spell-checker: disable
676        Ok(builder.append_operation(match &operation.op {
677            syn::BinOp::Add(_) => arith::addi(left, right, location),
678            syn::BinOp::Sub(_) => arith::subi(left, right, location),
679            syn::BinOp::Mul(_) => arith::muli(left, right, location),
680            syn::BinOp::Div(_) => arith::divsi(left, right, location),
681            syn::BinOp::Rem(_) => arith::remsi(left, right, location),
682            syn::BinOp::And(_) => arith::andi(left, right, location),
683            syn::BinOp::Or(_) => arith::ori(left, right, location),
684            syn::BinOp::BitXor(_) => arith::xori(left, right, location),
685            syn::BinOp::BitAnd(_) => arith::andi(left, right, location),
686            syn::BinOp::BitOr(_) => arith::ori(left, right, location),
687            syn::BinOp::Shl(_) => arith::shli(left, right, location),
688            syn::BinOp::Shr(_) => arith::shrsi(left, right, location),
689            syn::BinOp::Eq(_) => {
690                arith::cmpi(context, arith::CmpiPredicate::Eq, left, right, location)
691            }
692            syn::BinOp::Lt(_) => {
693                arith::cmpi(context, arith::CmpiPredicate::Slt, left, right, location)
694            }
695            syn::BinOp::Le(_) => {
696                arith::cmpi(context, arith::CmpiPredicate::Sle, left, right, location)
697            }
698            syn::BinOp::Ne(_) => {
699                arith::cmpi(context, arith::CmpiPredicate::Ne, left, right, location)
700            }
701            syn::BinOp::Ge(_) => {
702                arith::cmpi(context, arith::CmpiPredicate::Sge, left, right, location)
703            }
704            syn::BinOp::Gt(_) => {
705                arith::cmpi(context, arith::CmpiPredicate::Sgt, left, right, location)
706            }
707            syn::BinOp::AddAssign(_) => todo!(),
708            syn::BinOp::SubAssign(_) => todo!(),
709            syn::BinOp::MulAssign(_) => todo!(),
710            syn::BinOp::DivAssign(_) => todo!(),
711            syn::BinOp::RemAssign(_) => todo!(),
712            syn::BinOp::BitXorAssign(_) => todo!(),
713            syn::BinOp::BitAndAssign(_) => todo!(),
714            syn::BinOp::BitOrAssign(_) => todo!(),
715            syn::BinOp::ShlAssign(_) => todo!(),
716            syn::BinOp::ShrAssign(_) => todo!(),
717            _ => return Err(Error::NotSupported("unknown binary operator")),
718        }))
719        // spell-checker: enable
720    }
721
722    fn compile_expression_literal<'a>(
723        &self,
724        builder: &'a Block<'c>,
725        literal: &syn::ExprLit,
726    ) -> Result<OperationRef<'c, 'a>, Error> {
727        let context = self.context;
728        let location = Location::unknown(context);
729
730        Ok(builder.append_operation(match &literal.lit {
731            syn::Lit::Bool(boolean) => arith::constant(
732                context,
733                IntegerAttribute::new(IntegerType::new(context, 1).into(), boolean.value as i64)
734                    .into(),
735                location,
736            ),
737            syn::Lit::Char(_) => todo!(),
738            syn::Lit::Int(integer) => arith::constant(
739                context,
740                IntegerAttribute::new(
741                    match integer.suffix() {
742                        "" => Type::index(context),
743                        name => self.compile_primitive_type(name)?,
744                    },
745                    integer.base10_parse::<i64>()?,
746                )
747                .into(),
748                location,
749            ),
750            syn::Lit::Float(float) => arith::constant(
751                context,
752                FloatAttribute::new(
753                    context,
754                    match float.suffix() {
755                        "" => Type::index(context),
756                        name => self.compile_primitive_type(name)?,
757                    },
758                    float.base10_parse::<f64>()?,
759                )
760                .into(),
761                location,
762            ),
763            syn::Lit::Str(_) => todo!(),
764            syn::Lit::ByteStr(_) => todo!(),
765            syn::Lit::Byte(_) => todo!(),
766            _ => todo!(),
767        }))
768    }
769
770    fn compile_path<'a>(
771        &self,
772        builder: &'a Block<'c>,
773        path: &syn::ExprPath,
774        variables: &TrainMap<String, Value<'c, 'a>>,
775    ) -> Result<Value<'c, 'a>, Error> {
776        let context = self.context;
777        let name = self.convert_path_to_identifier(&path.path)?;
778
779        if let Some(&r#type) = self.functions.get(&name) {
780            Ok(builder
781                .append_operation(func::constant(
782                    context,
783                    FlatSymbolRefAttribute::new(context, &name),
784                    r#type,
785                    Location::unknown(context),
786                ))
787                .result(0)?
788                .into())
789        } else {
790            Ok(builder
791                .append_operation(memref::load(
792                    self.compile_variable(&name, variables)?,
793                    &[],
794                    Location::unknown(context),
795                ))
796                .result(0)?
797                .into())
798        }
799    }
800
801    fn convert_path_to_identifier(&self, path: &syn::Path) -> Result<String, Error> {
802        if let Some(identifier) = path.get_ident() {
803            Ok(identifier.to_string())
804        } else {
805            Err(Error::NotSupported("non-identifier path"))
806        }
807    }
808
809    fn compile_variable<'a>(
810        &self,
811        name: &str,
812        variables: &TrainMap<String, Value<'c, 'a>>,
813    ) -> Result<Value<'c, 'a>, Error> {
814        variables
815            .get(name)
816            .ok_or_else(|| Error::VariableNotDefined(name.into()))
817            .copied()
818    }
819
820    fn compile_unit<'a>(&self, builder: &'a Block<'c>) -> Result<Value<'c, 'a>, Error> {
821        let context = self.context;
822
823        Ok(builder
824            .append_operation(llvm::undef(
825                // TODO Should we use zero-field struct instead?
826                llvm::r#type::void(context),
827                Location::unknown(context),
828            ))
829            .result(0)?
830            .into())
831    }
832
833    fn get_struct_field_index(
834        &self,
835        member: &syn::Member,
836        info: &StructInfo,
837    ) -> Result<usize, Error> {
838        Ok(match member {
839            syn::Member::Named(name) => {
840                *info.field_indices.get(&name.to_string()).ok_or_else(|| {
841                    Error::StructFieldNotDefined(info.r#type.to_string(), name.to_string())
842                })?
843            }
844            syn::Member::Unnamed(index) => index.index as usize,
845        })
846    }
847
848    fn get_struct_info(&self, r#type: Type<'c>) -> Result<&StructInfo<'c>, Error> {
849        self.structs
850            .values()
851            .find(|info| info.r#type == r#type)
852            .ok_or_else(|| Error::StructNotDefined(r#type.to_string()))
853    }
854}
855
856#[cfg(test)]
857mod tests {
858    use super::*;
859    use crate::test::create_test_context;
860    use autophagy::math;
861    use melior::{ir::Location, Context};
862
863    fn compile<'c>(context: &'c Context, module: &Module<'c>, r#fn: &Fn) -> Result<(), Error> {
864        Compiler::new(context, module).compile_fn(r#fn)?;
865
866        Ok(())
867    }
868
869    #[test]
870    fn add() {
871        let context = create_test_context();
872
873        let location = Location::unknown(&context);
874        let module = Module::new(location);
875
876        compile(&context, &module, &math::add_fn()).unwrap();
877
878        assert!(module.as_operation().verify());
879    }
880
881    #[test]
882    fn sub() {
883        let context = create_test_context();
884
885        let location = Location::unknown(&context);
886        let module = Module::new(location);
887
888        compile(&context, &module, &math::sub_fn()).unwrap();
889
890        assert!(module.as_operation().verify());
891    }
892
893    #[test]
894    fn mul() {
895        let context = create_test_context();
896
897        let location = Location::unknown(&context);
898        let module = Module::new(location);
899
900        compile(&context, &module, &math::mul_fn()).unwrap();
901
902        assert!(module.as_operation().verify());
903    }
904
905    #[test]
906    fn div() {
907        let context = create_test_context();
908
909        let location = Location::unknown(&context);
910        let module = Module::new(location);
911
912        compile(&context, &module, &math::div_fn()).unwrap();
913
914        assert!(module.as_operation().verify());
915    }
916
917    #[test]
918    fn rem() {
919        let context = create_test_context();
920
921        let location = Location::unknown(&context);
922        let module = Module::new(location);
923
924        compile(&context, &module, &math::rem_fn()).unwrap();
925
926        assert!(module.as_operation().verify());
927    }
928
929    #[test]
930    fn neg() {
931        let context = create_test_context();
932
933        let location = Location::unknown(&context);
934        let module = Module::new(location);
935
936        compile(&context, &module, &math::neg_fn()).unwrap();
937
938        assert!(module.as_operation().verify());
939    }
940
941    #[test]
942    fn not() {
943        let context = create_test_context();
944
945        let location = Location::unknown(&context);
946        let module = Module::new(location);
947
948        compile(&context, &module, &math::not_fn()).unwrap();
949
950        assert!(module.as_operation().verify());
951    }
952
953    #[test]
954    fn and() {
955        let context = create_test_context();
956
957        let location = Location::unknown(&context);
958        let module = Module::new(location);
959
960        compile(&context, &module, &math::and_fn()).unwrap();
961
962        assert!(module.as_operation().verify());
963    }
964
965    #[test]
966    fn or() {
967        let context = create_test_context();
968
969        let location = Location::unknown(&context);
970        let module = Module::new(location);
971
972        compile(&context, &module, &math::or_fn()).unwrap();
973
974        assert!(module.as_operation().verify());
975    }
976
977    mod literal {
978        use super::*;
979
980        #[test]
981        fn bool() {
982            #[allow(dead_code)]
983            #[autophagy::quote]
984            fn foo() -> bool {
985                true
986            }
987
988            let context = create_test_context();
989
990            let location = Location::unknown(&context);
991            let module = Module::new(location);
992
993            compile(&context, &module, &foo_fn()).unwrap();
994
995            assert!(module.as_operation().verify());
996        }
997
998        #[test]
999        fn float32() {
1000            #[allow(dead_code)]
1001            #[autophagy::quote]
1002            fn foo() -> f32 {
1003                42f32
1004            }
1005
1006            let context = create_test_context();
1007
1008            let location = Location::unknown(&context);
1009            let module = Module::new(location);
1010
1011            compile(&context, &module, &foo_fn()).unwrap();
1012
1013            assert!(module.as_operation().verify());
1014        }
1015
1016        #[test]
1017        fn float64() {
1018            #[allow(dead_code)]
1019            #[autophagy::quote]
1020            fn foo() -> f64 {
1021                42f64
1022            }
1023
1024            let context = create_test_context();
1025
1026            let location = Location::unknown(&context);
1027            let module = Module::new(location);
1028
1029            compile(&context, &module, &foo_fn()).unwrap();
1030
1031            assert!(module.as_operation().verify());
1032        }
1033    }
1034
1035    #[test]
1036    fn call() {
1037        #[allow(dead_code)]
1038        #[autophagy::quote]
1039        fn foo(x: usize, y: usize) -> usize {
1040            x + y
1041        }
1042
1043        #[allow(dead_code)]
1044        #[autophagy::quote]
1045        fn bar() -> usize {
1046            foo(1, 2)
1047        }
1048
1049        let context = create_test_context();
1050
1051        let location = Location::unknown(&context);
1052        let module = Module::new(location);
1053
1054        let mut compiler = Compiler::new(&context, &module);
1055
1056        compiler.compile_fn(&foo_fn()).unwrap();
1057        compiler.compile_fn(&bar_fn()).unwrap();
1058
1059        assert!(module.as_operation().verify());
1060    }
1061
1062    #[test]
1063    fn dereference() {
1064        #[allow(dead_code)]
1065        #[autophagy::quote]
1066        fn foo(x: &usize) -> usize {
1067            *x
1068        }
1069
1070        let context = create_test_context();
1071
1072        let location = Location::unknown(&context);
1073        let module = Module::new(location);
1074
1075        compile(&context, &module, &foo_fn()).unwrap();
1076
1077        assert!(module.as_operation().verify());
1078    }
1079
1080    #[test]
1081    fn r#if() {
1082        #[allow(dead_code)]
1083        #[autophagy::quote]
1084        fn foo() -> usize {
1085            if true {
1086                42usize
1087            } else {
1088                13usize
1089            }
1090        }
1091
1092        let context = create_test_context();
1093
1094        let location = Location::unknown(&context);
1095        let module = Module::new(location);
1096
1097        compile(&context, &module, &foo_fn()).unwrap();
1098
1099        assert!(module.as_operation().verify());
1100    }
1101
1102    #[test]
1103    fn r#let() {
1104        #[allow(dead_code, clippy::let_and_return)]
1105        #[autophagy::quote]
1106        fn foo() -> usize {
1107            let x = 42usize;
1108
1109            x
1110        }
1111
1112        let context = create_test_context();
1113
1114        let location = Location::unknown(&context);
1115        let module = Module::new(location);
1116
1117        compile(&context, &module, &foo_fn()).unwrap();
1118
1119        assert!(module.as_operation().verify());
1120    }
1121
1122    #[test]
1123    fn r#loop() {
1124        #[allow(dead_code)]
1125        #[autophagy::quote]
1126        fn foo() {
1127            #[allow(clippy::empty_loop)]
1128            loop {}
1129        }
1130
1131        let context = create_test_context();
1132
1133        let location = Location::unknown(&context);
1134        let module = Module::new(location);
1135
1136        compile(&context, &module, &foo_fn()).unwrap();
1137
1138        assert!(module.as_operation().verify());
1139    }
1140
1141    #[test]
1142    fn struct_field() {
1143        #[autophagy::quote]
1144        struct Foo {
1145            bar: i32,
1146        }
1147
1148        #[allow(dead_code)]
1149        #[autophagy::quote]
1150        fn foo(x: Foo) -> i32 {
1151            x.bar
1152        }
1153
1154        let context = create_test_context();
1155
1156        let location = Location::unknown(&context);
1157        let module = Module::new(location);
1158        let mut compiler = Compiler::new(&context, &module);
1159
1160        compiler.compile_struct(&foo_struct()).unwrap();
1161        compiler.compile_fn(&foo_fn()).unwrap();
1162
1163        assert!(module.as_operation().verify());
1164    }
1165
1166    #[test]
1167    fn struct_field_dereference() {
1168        #[autophagy::quote]
1169        struct Foo {
1170            bar: i32,
1171        }
1172
1173        #[allow(dead_code)]
1174        #[autophagy::quote]
1175        fn foo(x: &Foo) -> i32 {
1176            x.bar
1177        }
1178
1179        let context = create_test_context();
1180
1181        let location = Location::unknown(&context);
1182        let module = Module::new(location);
1183        let mut compiler = Compiler::new(&context, &module);
1184
1185        compiler.compile_struct(&foo_struct()).unwrap();
1186        compiler.compile_fn(&foo_fn()).unwrap();
1187
1188        assert!(module.as_operation().verify());
1189    }
1190
1191    #[test]
1192    fn struct_field_double_dereference() {
1193        #[autophagy::quote]
1194        struct Foo {
1195            bar: i32,
1196        }
1197
1198        #[allow(dead_code)]
1199        #[autophagy::quote]
1200        fn foo(x: &&Foo) -> i32 {
1201            x.bar
1202        }
1203
1204        let context = create_test_context();
1205
1206        let location = Location::unknown(&context);
1207        let module = Module::new(location);
1208        let mut compiler = Compiler::new(&context, &module);
1209
1210        compiler.compile_struct(&foo_struct()).unwrap();
1211        compiler.compile_fn(&foo_fn()).unwrap();
1212
1213        assert!(module.as_operation().verify());
1214    }
1215
1216    #[test]
1217    fn struct_field_assignment() {
1218        #[autophagy::quote]
1219        struct Foo {
1220            bar: i32,
1221        }
1222
1223        #[allow(dead_code)]
1224        #[autophagy::quote]
1225        fn foo(x: &mut Foo) {
1226            x.bar = 42i32;
1227        }
1228
1229        let context = create_test_context();
1230
1231        let location = Location::unknown(&context);
1232        let module = Module::new(location);
1233        let mut compiler = Compiler::new(&context, &module);
1234
1235        compiler.compile_struct(&foo_struct()).unwrap();
1236        compiler.compile_fn(&foo_fn()).unwrap();
1237
1238        assert!(module.as_operation().verify());
1239    }
1240
1241    #[test]
1242    fn struct_literal() {
1243        #[allow(dead_code)]
1244        #[autophagy::quote]
1245        struct Foo {
1246            bar: i32,
1247            baz: f64,
1248        }
1249
1250        #[allow(dead_code)]
1251        #[autophagy::quote]
1252        fn foo() -> Foo {
1253            Foo {
1254                bar: 42i32,
1255                baz: 1.5f64,
1256            }
1257        }
1258
1259        let context = create_test_context();
1260
1261        let location = Location::unknown(&context);
1262        let module = Module::new(location);
1263        let mut compiler = Compiler::new(&context, &module);
1264
1265        compiler.compile_struct(&foo_struct()).unwrap();
1266        compiler.compile_fn(&foo_fn()).unwrap();
1267
1268        assert!(module.as_operation().verify());
1269    }
1270
1271    #[test]
1272    fn r#while() {
1273        #[allow(dead_code)]
1274        #[autophagy::quote]
1275        fn foo() {
1276            #[allow(while_true)]
1277            while true {}
1278        }
1279
1280        let context = create_test_context();
1281
1282        let location = Location::unknown(&context);
1283        let module = Module::new(location);
1284
1285        compile(&context, &module, &foo_fn()).unwrap();
1286
1287        assert!(module.as_operation().verify());
1288    }
1289}