Skip to main content

citadel_backend/asm/codegen/
mod.rs

1//! This is the compiler for translating the IR to assembly
2//! Future: This will use the low-level IR at some point but
3//!         until the lir is finished, it will use the high-level IR
4//!
5//! Generally this is only serves as a helper for the actual Backend#compile
6//! function.
7
8use std::collections::{HashMap, HashSet};
9
10use citadel_frontend::{
11    ir::{
12        self, irgen::TypeTable, ArithOpExpr, BlockStmt, CallExpr, ExitStmt, FuncStmt, IRExpr, IRStmt, Ident, JumpStmt, LabelStmt, ReturnStmt, StructInitExpr, Type, VarStmt, INT16_T, INT32_T, INT64_T, INT8_T
13    },
14    util::CompositeDataType,
15};
16
17use crate::asm::{
18    elements::{
19        AsmElement, DataSize, Declaration, Directive, DirectiveType, Instruction, Label, Literal,
20        Opcode, Operand, Register, Size, SizedLiteral, StdFunction,
21    },
22    utils::codegen as cutils,
23};
24
25pub const FUNCTION_ARG_REGISTERS_8: [Register; 6] = [
26    Register::Al,
27    Register::Bl,
28    Register::Cl,
29    Register::Dl,
30    Register::R9b,
31    Register::R10b,
32];
33
34pub const FUNCTION_ARG_REGISTERS_16: [Register; 6] = [
35    Register::Ax,
36    Register::Bx,
37    Register::Cx,
38    Register::Dx,
39    Register::R9w,
40    Register::R10w,
41];
42
43pub const FUNCTION_ARG_REGISTERS_32: [Register; 6] = [
44    Register::Edi,
45    Register::Esi,
46    Register::Edx,
47    Register::Ecx,
48    Register::R9d,
49    Register::R10d,
50];
51
52pub const FUNCTION_ARG_REGISTERS_64: [Register; 6] = [
53    Register::Rdi,
54    Register::Rsi,
55    Register::Rdx,
56    Register::Rcx,
57    Register::R9,
58    Register::R10,
59];
60
61#[derive(Default)]
62pub struct CodeGenerator<'c> {
63    pub out: Vec<AsmElement>,
64    pub types: TypeTable<'c>,
65
66    // Literals
67    /// Read only data section
68    pub rodata: Vec<Declaration>,
69    pub data: Vec<Declaration>,
70    /// Literal constant index
71    pub lc_index: usize,
72
73    pub defined_functions: HashSet<StdFunction>,
74    pub symbol_table: HashMap<&'c str, i32>,
75
76    pub stack_pointer: i32,
77}
78
79impl<'c> CodeGenerator<'c> {
80    pub fn new(types: TypeTable<'c>) -> Self {
81        Self {
82            types,
83            ..Default::default()
84        }
85    }
86
87    pub fn gen_stmt(&mut self, node: &'c IRStmt) {
88        match node {
89            IRStmt::DeclaredFunction(_) => todo!(),
90            IRStmt::Function(node) => self.gen_function(node),
91            IRStmt::Entry(node) => self.gen_entry(node),
92            IRStmt::Struct(_) => (),
93            IRStmt::Union(_) => (),
94            IRStmt::Variable(node) => self.gen_variable(node),
95            IRStmt::Label(node) => self.gen_label(node),
96            IRStmt::Return(node) => self.gen_return(node),
97            IRStmt::Exit(node) => self.gen_exit(node),
98            IRStmt::Jump(node) => self.gen_jump(node),
99            IRStmt::Call(node) => self.gen_call(node),
100        }
101    }
102
103    fn gen_expr(&mut self, node: &'c IRExpr) -> Operand {
104        match &node {
105            IRExpr::Literal(node, type_) => match node {
106                ir::Literal::Int32(val) => Operand::SizedLiteral(SizedLiteral(Literal::Int32(*val), DataSize::DWord)),
107                ir::Literal::String(val) => self.gen_string(val, type_),
108                int => todo!("Handle non-i32 literals here: {:?}", int),
109            },
110            IRExpr::Call(node) => {
111                self.gen_call(node);
112                let reg = Register::Rax;
113                Operand::Register(reg)
114            }
115            IRExpr::ArithOp(node) => self.gen_arith_op(node, true),
116            IRExpr::Ident(node) => cutils::get_stack_location(
117                *self
118                    .symbol_table
119                    .get(node)
120                    .unwrap_or_else(|| panic!("Could not find ident with name {node:?}")),
121            ),
122            IRExpr::StructInit(node) => self.gen_struct_init(node),
123        }
124    }
125
126    pub fn gen_entry(&mut self, node: &'c BlockStmt<'c>) {
127        // Text directive (entry point)
128        self.out.push(AsmElement::Directive(Directive {
129            _type: DirectiveType::Text,
130        }));
131        self.out.push(AsmElement::Declaration(Declaration::Global(
132            "_start".to_string(),
133        )));
134
135        // _start label
136        self.out.push(AsmElement::Label(Label {
137            name: "_start".to_string(),
138        }));
139        for stmt in &node.stmts {
140            self.gen_stmt(stmt);
141        }
142    }
143
144    fn gen_call(&mut self, node: &'c CallExpr) {
145        match node.name {
146            "print" => self.gen_print(node),
147            _ => {
148                self.gen_call_args(node);
149                self.out.push(cutils::gen_call(&node.name))
150            }
151        }
152    }
153
154    fn gen_jump(&mut self, node: &'c JumpStmt) {
155        self.out.push(AsmElement::Instruction(Instruction {
156            opcode: Opcode::Jmp,
157            args: vec![Operand::Ident(node.label.to_string())],
158        }))
159    }
160
161    fn gen_string(&mut self, val: &str, type_: &Type<'c>) -> Operand {
162        let size = *match type_ {
163            Type::Ident(_) => todo!(),
164            Type::Array(_, len) => len,
165        };
166        // TODO: use different splitting techniques based on string length
167        let mut strings = cutils::split_string(val, 8);
168        let last_string = strings.pop().unwrap();
169        self.stack_pointer -= size as i32;
170        Operand::SizedLiteral(SizedLiteral(
171            Literal::Int64(cutils::conv_str_to_bytes(last_string) as i64),
172            cutils::word_from_size(size as u8),
173        ))
174    }
175
176    fn gen_arith_op(&mut self, node: &'c ArithOpExpr, move_to_rax: bool) -> Operand {
177        if move_to_rax {
178            let left_expr = self.gen_expr(&node.values.0);
179            self.gen_mov_ins(Operand::Register(Register::Rax), left_expr)
180        }
181        let arith_op = match node.op {
182            ir::Operator::Add => self.gen_arith_op_ins(Opcode::Add, node),
183            ir::Operator::Sub => self.gen_arith_op_ins(Opcode::Sub, node),
184            ir::Operator::Mul => self.gen_arith_op_ins(Opcode::Mul, node),
185            ir::Operator::Div => self.gen_arith_op_ins(Opcode::Div, node),
186        };
187        self.out.push(arith_op);
188        Operand::Register(Register::Rax)
189    }
190
191    fn gen_arith_op_ins(&mut self, opcode: Opcode, node: &'c ArithOpExpr) -> AsmElement {
192        AsmElement::Instruction(Instruction {
193            opcode,
194            args: vec![
195                Operand::Register(Register::Rax),
196                match &*node.values.1 {
197                    IRExpr::ArithOp(expr) => self.gen_arith_op(expr, false),
198                    expr => self.gen_expr(expr),
199                },
200            ],
201        })
202    }
203
204    fn gen_return(&mut self, node: &'c ReturnStmt) {
205        let val = self.gen_expr(&node.ret_val);
206        self.out
207            .push(cutils::gen_mov_ins(Operand::Register(Register::Rax), val));
208        self.out.push(cutils::destroy_stackframe());
209        self.out.push(cutils::gen_ret());
210    }
211
212    fn gen_exit(&mut self, node: &'c ExitStmt) {
213        let expr = self.gen_expr(&node.exit_code);
214        self.out
215            .push(cutils::gen_mov_ins(Operand::Register(Register::Rdi), expr));
216        self.gen_mov_ins(
217            Operand::Register(Register::Rax),
218            Operand::Literal(Literal::Int32(60)),
219        );
220        self.out.push(cutils::gen_syscall());
221    }
222
223    fn gen_variable(&mut self, node: &'c VarStmt) {
224        let size = self.size_of(&node.name._type);
225        let mut val = self.gen_expr(&node.val);
226        // FIXME: This is a hack to ensure that the size does not get decremented for arrays
227        if let Type::Ident(_) = node.name._type {
228            self.stack_pointer -= size as i32
229        }
230
231        if let Operand::Literal(lit) = val {
232            val = Operand::SizedLiteral(cutils::literal_to_sized_literal(lit)
233                .expect("Failed to convert literal to sized literal, most likely caused due to usage of float which are not supported yet"))
234        };
235
236        if let Operand::SizedLiteral(SizedLiteral(lit, DataSize::QWord)) = val {
237            self.gen_mov_ins(Operand::Register(Register::Rax), Operand::Literal(lit));
238            val = Operand::Register(Register::Rax);
239        }
240
241        self.gen_mov_ins(cutils::get_stack_location(self.stack_pointer), val);
242
243        self.symbol_table
244            .insert(&node.name.ident, self.stack_pointer);
245    }
246
247    fn gen_function(&mut self, node: &'c FuncStmt) {
248        self.out.push(AsmElement::Label(Label {
249            name: node.name.ident.to_string(),
250        }));
251
252        let stack_frame = cutils::create_stackframe();
253
254        self.out.push(stack_frame.0);
255        self.out.push(stack_frame.1);
256
257        self.gen_args(node);
258
259        for stmt in &node.block.stmts {
260            self.gen_stmt(stmt);
261        }
262
263        if let Some(elem) = self.out.last() {
264            match elem {
265                AsmElement::Instruction(Instruction {
266                    opcode: Opcode::Ret,
267                    args: _,
268                }) => (),
269                _ => {
270                    if let ir::Type::Ident("void") = node.name._type {
271                        self.out.push(cutils::destroy_stackframe());
272                    }
273                    self.out.push(cutils::gen_ret());
274                }
275            }
276        }
277    }
278
279    fn gen_struct_init(&mut self, node: &'c StructInitExpr) -> Operand {
280        let size = self.size_of(&ir::Type::Ident(node.name));
281        self.gen_mov_ins(
282            cutils::get_stack_location(self.stack_pointer - size as i32),
283            Operand::Literal(Literal::Int32(0)),
284        );
285        self.stack_pointer -= size as i32;
286        // TODO: Use type suffixes for this
287        for (i, val) in node.values.iter().enumerate() {
288            let fields = &self.types.get(&node.name).unwrap().1;
289            let _field = &fields[i];
290            let expr = self.gen_expr(val);
291            self.gen_mov_ins(cutils::get_stack_location(0), expr);
292        }
293        todo!()
294    }
295
296    fn gen_args(&mut self, node: &'c FuncStmt) {
297        for (i, expr) in node.args.iter().enumerate() {
298            let size = self.size_of(&expr._type);
299            self.gen_mov_ins(
300                cutils::get_stack_location(self.stack_pointer - size as i32),
301                Operand::Register(cutils::arg_regs_by_size(size.try_into().expect("Failed to convert u32 to u8"))[i]),
302            );
303            self.stack_pointer -= size as i32;
304            self.symbol_table.insert(&expr.ident, self.stack_pointer);
305        }
306    }
307
308    fn gen_call_args(&mut self, node: &'c CallExpr) {
309        for (i, expr) in node.args.iter().enumerate() {
310            let val = self.gen_expr(expr);
311            self.gen_mov_ins(
312                Operand::Register(cutils::arg_regs_by_size(val.size())[i]),
313                val,
314            );
315        }
316    }
317
318    fn gen_print(&mut self, node: &'c CallExpr) {
319        let arg = self.gen_expr(
320            node.args
321                .first()
322                .expect("Print function neeeds at least one argument"),
323        );
324        self.gen_mov_ins(Operand::Register(Register::Rsi), arg);
325        self.gen_mov_ins(
326            Operand::Register(Register::Rdx),
327            Operand::Literal(Literal::Int8(8)),
328        );
329        self.out.push(cutils::gen_call("print"));
330        self.defined_functions.insert(StdFunction::Print);
331    }
332
333    fn gen_label(&mut self, node: &'c LabelStmt) {
334        self.out.push(AsmElement::Label(Label {
335            name: node.name.to_string(),
336        }));
337    }
338
339    /// Returns the size of the type in bytes
340    fn size_of(&self, _type: &ir::Type<'c>) -> u32 {
341        // The type or array is an integer type/array
342        match _type {
343            Type::Ident(ident @ (INT8_T | INT16_T | INT32_T | INT64_T)) => {
344                return cutils::int_size(*ident) as u32;
345            }
346            Type::Array(Type::Ident(ident @ (INT8_T | INT16_T | INT32_T | INT64_T)), size) => {
347                return cutils::int_size(*ident) as u32 * *size;
348            }
349            _ => (),
350        }
351
352        let type_name = match _type {
353            Type::Ident(ident) => ident,
354            Type::Array(ident, _) => match ident {
355                Type::Ident(id) => id,
356                Type::Array(id, _) => return self.size_of(id),
357            },
358        };
359
360        let cdt = self
361            .types
362            .get(type_name)
363            .unwrap_or_else(|| panic!("Could not find type with the name {}", _type));
364        let mut size: u32 = 0;
365        match cdt.0 {
366            // Add sizes if cdt is a struct
367            CompositeDataType::Struct => {
368                for field in &cdt.1 {
369                    size += self.size_of(&field._type);
370                }
371            }
372            // Use largest size if cdt is a union
373            CompositeDataType::Union => {
374                for variant in &cdt.1 {
375                    let size1 = self.size_of(&variant._type);
376                    if size1 > size {
377                        size = size1;
378                    }
379                }
380            }
381        }
382        match _type {
383            Type::Ident(_) => size,
384            Type::Array(_, arr_size) => size * *arr_size,
385        }
386    }
387
388    fn gen_mov_ins(&mut self, target: Operand, val: Operand) {
389        if target != val {
390            self.out.push(cutils::gen_mov_ins(target, val))
391        }
392    }
393}