1use 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 pub rodata: Vec<Declaration>,
69 pub data: Vec<Declaration>,
70 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 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 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 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 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 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 fn size_of(&self, _type: &ir::Type<'c>) -> u32 {
341 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 CompositeDataType::Struct => {
368 for field in &cdt.1 {
369 size += self.size_of(&field._type);
370 }
371 }
372 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}