Skip to main content

palladium/codegen/
llvm_backend.rs

1// LLVM backend for Palladium
2// "From Turing's proofs to von Neumann's performance"
3
4use crate::ast::{Expr, Function, Item, Program, Stmt, Type};
5use crate::errors::{CompileError, Result};
6use std::path::PathBuf;
7
8/// LLVM backend code generator
9pub struct LLVMCodeGenerator {
10    module_name: String,
11    /// Whether LLVM is available
12    llvm_available: bool,
13}
14
15impl LLVMCodeGenerator {
16    pub fn new(module_name: &str) -> Result<Self> {
17        Ok(Self {
18            module_name: module_name.to_string(),
19            llvm_available: Self::check_llvm_availability(),
20        })
21    }
22
23    /// Check if LLVM is available on the system
24    fn check_llvm_availability() -> bool {
25        // For now, we'll check if LLVM is available via llvm-config
26        std::process::Command::new("llvm-config")
27            .arg("--version")
28            .output()
29            .is_ok()
30    }
31
32    /// Compile a program to LLVM IR
33    pub fn compile(&mut self, program: &Program) -> Result<()> {
34        if !self.llvm_available {
35            println!("   Warning: LLVM tools not found. Generating LLVM IR text only.");
36        }
37
38        // For now, generate LLVM IR as text
39        let ir = self.generate_ir(program)?;
40
41        // Write to .ll file
42        let output_path = PathBuf::from("build_output").join(format!("{}.ll", self.module_name));
43        std::fs::write(&output_path, ir)?;
44
45        println!("   Generated LLVM IR: {}", output_path.display());
46
47        Ok(())
48    }
49
50    /// Generate LLVM IR for the program
51    fn generate_ir(&self, program: &Program) -> Result<String> {
52        let mut ir = String::new();
53
54        // Module header
55        ir.push_str(&format!("; ModuleID = '{}'\n", self.module_name));
56        ir.push_str("source_filename = \"palladium\"\n");
57        ir.push_str("target datalayout = \"e-m:e-p270:32:32-p271:32:32-p272:64:64-i64:64-f80:128-n8:16:32:64-S128\"\n");
58        ir.push_str("target triple = \"x86_64-pc-linux-gnu\"\n\n");
59
60        // External function declarations
61        ir.push_str("; External function declarations\n");
62        ir.push_str("declare i32 @printf(i8*, ...)\n");
63        ir.push_str("declare i8* @malloc(i64)\n");
64        ir.push_str("declare void @free(i8*)\n");
65        ir.push_str("declare i64 @strlen(i8*)\n");
66        ir.push_str("declare i8* @strcpy(i8*, i8*)\n");
67        ir.push_str("declare i8* @strcat(i8*, i8*)\n");
68        ir.push_str("declare i32 @strcmp(i8*, i8*)\n\n");
69
70        // String constants for print functions
71        ir.push_str("; String constants\n");
72        ir.push_str(
73            "@.str_fmt = private unnamed_addr constant [4 x i8] c\"%s\\0A\\00\", align 1\n",
74        );
75        ir.push_str(
76            "@.int_fmt = private unnamed_addr constant [6 x i8] c\"%lld\\0A\\00\", align 1\n\n",
77        );
78
79        // Generate functions
80        for item in &program.items {
81            match item {
82                Item::Function(func) => {
83                    ir.push_str(&self.generate_function(func)?);
84                    ir.push('\n');
85                }
86                _ => {
87                    // Skip other items for now
88                }
89            }
90        }
91
92        Ok(ir)
93    }
94
95    /// Generate LLVM IR for a function
96    fn generate_function(&self, func: &Function) -> Result<String> {
97        let mut ir = String::new();
98
99        // Function signature
100        let ret_type = self.type_to_llvm(&func.return_type);
101        ir.push_str(&format!("define {} @{}(", ret_type, func.name));
102
103        // Parameters
104        for (i, param) in func.params.iter().enumerate() {
105            if i > 0 {
106                ir.push_str(", ");
107            }
108            let param_type = self.type_to_llvm(&Some(param.ty.clone()));
109            ir.push_str(&format!("{} %{}", param_type, param.name));
110        }
111
112        ir.push_str(") {\n");
113        ir.push_str("entry:\n");
114
115        // Function body
116        let mut label_counter = 0;
117        let mut var_counter = 0;
118
119        for stmt in &func.body {
120            ir.push_str(&self.generate_statement(stmt, &mut var_counter, &mut label_counter)?);
121        }
122
123        // Default return if needed
124        if func.return_type.is_none() && !func.body.iter().any(|s| matches!(s, Stmt::Return(_))) {
125            ir.push_str("  ret void\n");
126        }
127
128        ir.push_str("}\n");
129
130        Ok(ir)
131    }
132
133    /// Convert Palladium type to LLVM type
134    fn type_to_llvm(&self, ty: &Option<Type>) -> &'static str {
135        match ty {
136            None => "void",
137            Some(Type::I32) => "i32",
138            Some(Type::I64) => "i64",
139            Some(Type::U32) => "i32",
140            Some(Type::U64) => "i64",
141            Some(Type::Bool) => "i1",
142            Some(Type::String) => "i8*",
143            Some(Type::Unit) => "void",
144            _ => "i8*", // Default to pointer for complex types
145        }
146    }
147
148    /// Generate LLVM IR for a statement
149    fn generate_statement(
150        &self,
151        stmt: &Stmt,
152        var_counter: &mut i32,
153        label_counter: &mut i32,
154    ) -> Result<String> {
155        let mut ir = String::new();
156
157        match stmt {
158            Stmt::Expr(expr) => {
159                let (expr_ir, _) = self.generate_expression(expr, var_counter)?;
160                ir.push_str(&expr_ir);
161            }
162
163            Stmt::Let { name, value, .. } => {
164                let (expr_ir, result_var) = self.generate_expression(value, var_counter)?;
165                ir.push_str(&expr_ir);
166
167                // For simplicity, we'll use the variable name directly
168                // In a real implementation, we'd need proper SSA form
169                ir.push_str(&format!("  %{} = alloca i64\n", name));
170                ir.push_str(&format!("  store i64 {}, i64* %{}\n", result_var, name));
171            }
172
173            Stmt::Return(Some(expr)) => {
174                let (expr_ir, result) = self.generate_expression(expr, var_counter)?;
175                ir.push_str(&expr_ir);
176                ir.push_str(&format!("  ret i64 {}\n", result));
177            }
178
179            Stmt::Return(None) => {
180                ir.push_str("  ret void\n");
181            }
182
183            Stmt::If {
184                condition,
185                then_branch,
186                else_branch,
187                ..
188            } => {
189                let then_label = format!("then{}", label_counter);
190                let else_label = format!("else{}", label_counter);
191                let end_label = format!("endif{}", label_counter);
192                *label_counter += 1;
193
194                let (cond_ir, cond_result) = self.generate_expression(condition, var_counter)?;
195                ir.push_str(&cond_ir);
196
197                if else_branch.is_some() {
198                    ir.push_str(&format!(
199                        "  br i1 {}, label %{}, label %{}\n",
200                        cond_result, then_label, else_label
201                    ));
202                } else {
203                    ir.push_str(&format!(
204                        "  br i1 {}, label %{}, label %{}\n",
205                        cond_result, then_label, end_label
206                    ));
207                }
208
209                // Then branch
210                ir.push_str(&format!("{}:\n", then_label));
211                for stmt in then_branch {
212                    ir.push_str(&self.generate_statement(stmt, var_counter, label_counter)?);
213                }
214                ir.push_str(&format!("  br label %{}\n", end_label));
215
216                // Else branch
217                if let Some(else_stmts) = else_branch {
218                    ir.push_str(&format!("{}:\n", else_label));
219                    for stmt in else_stmts {
220                        ir.push_str(&self.generate_statement(stmt, var_counter, label_counter)?);
221                    }
222                    ir.push_str(&format!("  br label %{}\n", end_label));
223                }
224
225                // End label
226                ir.push_str(&format!("{}:\n", end_label));
227            }
228
229            _ => {
230                // TODO: Implement other statements
231                ir.push_str("  ; TODO: Implement this statement\n");
232            }
233        }
234
235        Ok(ir)
236    }
237
238    /// Generate LLVM IR for an expression
239    /// Returns (IR code, result variable/value)
240    #[allow(clippy::only_used_in_recursion)]
241    fn generate_expression(&self, expr: &Expr, var_counter: &mut i32) -> Result<(String, String)> {
242        let mut ir = String::new();
243
244        match expr {
245            Expr::Integer(n) => Ok((String::new(), n.to_string())),
246
247            Expr::Bool(b) => Ok((String::new(), if *b { "1" } else { "0" }.to_string())),
248
249            Expr::String(s) => {
250                // Create a string constant
251                let const_name = format!("@.str.{}", var_counter);
252                *var_counter += 1;
253
254                let escaped = s
255                    .replace("\\", "\\\\")
256                    .replace("\"", "\\\"")
257                    .replace("\n", "\\n");
258                ir.push_str(&format!(
259                    "{} = private unnamed_addr constant [{} x i8] c\"{}\\00\"\n",
260                    const_name,
261                    s.len() + 1,
262                    escaped
263                ));
264
265                let ptr_var = format!("%str.{}", var_counter);
266                *var_counter += 1;
267                ir.push_str(&format!(
268                    "  {} = getelementptr [{} x i8], [{} x i8]* {}, i32 0, i32 0\n",
269                    ptr_var,
270                    s.len() + 1,
271                    s.len() + 1,
272                    const_name
273                ));
274
275                Ok((ir, ptr_var))
276            }
277
278            Expr::Ident(name) => {
279                let var = format!("%{}.load", var_counter);
280                *var_counter += 1;
281                ir.push_str(&format!("  {} = load i64, i64* %{}\n", var, name));
282                Ok((ir, var))
283            }
284
285            Expr::Binary {
286                left, op, right, ..
287            } => {
288                let (left_ir, left_var) = self.generate_expression(left, var_counter)?;
289                let (right_ir, right_var) = self.generate_expression(right, var_counter)?;
290
291                ir.push_str(&left_ir);
292                ir.push_str(&right_ir);
293
294                let result_var = format!("%{}", var_counter);
295                *var_counter += 1;
296
297                let op_str = match op {
298                    crate::ast::BinOp::Add => "add",
299                    crate::ast::BinOp::Sub => "sub",
300                    crate::ast::BinOp::Mul => "mul",
301                    crate::ast::BinOp::Div => "sdiv",
302                    crate::ast::BinOp::Mod => "srem",
303                    crate::ast::BinOp::Lt => "icmp slt",
304                    crate::ast::BinOp::Le => "icmp sle",
305                    crate::ast::BinOp::Gt => "icmp sgt",
306                    crate::ast::BinOp::Ge => "icmp sge",
307                    crate::ast::BinOp::Eq => "icmp eq",
308                    crate::ast::BinOp::Ne => "icmp ne",
309                    _ => {
310                        return Err(CompileError::Generic(
311                            "Unsupported binary operator".to_string(),
312                        ))
313                    }
314                };
315
316                ir.push_str(&format!(
317                    "  {} = {} i64 {}, {}\n",
318                    result_var, op_str, left_var, right_var
319                ));
320
321                Ok((ir, result_var))
322            }
323
324            Expr::Call { func, args, .. } => {
325                if let Expr::Ident(func_name) = func.as_ref() {
326                    match func_name.as_str() {
327                        "print" => {
328                            if args.len() == 1 {
329                                let (arg_ir, arg_var) =
330                                    self.generate_expression(&args[0], var_counter)?;
331                                ir.push_str(&arg_ir);
332                                ir.push_str(&format!("  call i32 (i8*, ...) @printf(i8* getelementptr inbounds ([4 x i8], [4 x i8]* @.str_fmt, i32 0, i32 0), i8* {})\n", arg_var));
333                            }
334                            Ok((ir, "%0".to_string())) // Dummy return
335                        }
336                        "print_int" => {
337                            if args.len() == 1 {
338                                let (arg_ir, arg_var) =
339                                    self.generate_expression(&args[0], var_counter)?;
340                                ir.push_str(&arg_ir);
341                                ir.push_str(&format!("  call i32 (i8*, ...) @printf(i8* getelementptr inbounds ([6 x i8], [6 x i8]* @.int_fmt, i32 0, i32 0), i64 {})\n", arg_var));
342                            }
343                            Ok((ir, "%0".to_string())) // Dummy return
344                        }
345                        _ => {
346                            // User-defined function call
347                            let mut arg_vars = Vec::new();
348                            for arg in args {
349                                let (arg_ir, arg_var) =
350                                    self.generate_expression(arg, var_counter)?;
351                                ir.push_str(&arg_ir);
352                                arg_vars.push(arg_var);
353                            }
354
355                            let result_var = format!("%{}", var_counter);
356                            *var_counter += 1;
357
358                            ir.push_str(&format!("  {} = call i64 @{}(", result_var, func_name));
359                            for (i, arg_var) in arg_vars.iter().enumerate() {
360                                if i > 0 {
361                                    ir.push_str(", ");
362                                }
363                                ir.push_str(&format!("i64 {}", arg_var));
364                            }
365                            ir.push_str(")\n");
366
367                            Ok((ir, result_var))
368                        }
369                    }
370                } else {
371                    Err(CompileError::Generic(
372                        "Complex function calls not yet supported".to_string(),
373                    ))
374                }
375            }
376
377            _ => {
378                // TODO: Implement other expressions
379                Ok((String::new(), "%0".to_string()))
380            }
381        }
382    }
383
384    /// Write the generated LLVM IR to a file
385    pub fn write_output(&self) -> Result<PathBuf> {
386        let build_dir = PathBuf::from("build_output");
387        if !build_dir.exists() {
388            std::fs::create_dir_all(&build_dir)?;
389        }
390
391        let output_path = build_dir.join(format!("{}.ll", self.module_name));
392
393        Ok(output_path)
394    }
395}
396
397#[cfg(test)]
398mod tests {
399    use super::*;
400
401    #[test]
402    fn test_llvm_availability() {
403        let available = LLVMCodeGenerator::check_llvm_availability();
404        println!("LLVM available: {}", available);
405    }
406}