refactor(sema): Separated from ir

This commit is contained in:
2026-05-31 20:59:18 +08:00
parent 1eca4a225b
commit c42575c1c6
10 changed files with 829 additions and 260 deletions
+449
View File
@@ -0,0 +1,449 @@
use std::collections::BTreeMap;
use crate::{
ast::types::{
BinaryOp, BlockStmt, BreakStmt, CompileUnit, ContinueStmt, Expr, ExprValue, FuncDeclStmt,
GlobalDeclStmt, IfElseBranch, IfStmt, ReturnStmt, Statement, VarDeclStmt, WhileStmt,
},
diagnostic::{span::Span, Diagnositics},
sema::{
err::SemaError,
hir::{
HirBlockStmt, HirBreakStmt, HirCompileUnit, HirContinueStmt, HirExpr, HirExprValue,
HirFuncDeclStmt, HirGlobalDeclStmt, HirIfElseBranch, HirIfStmt, HirParam,
HirReturnStmt, HirStatement, HirVarDeclStmt, HirVarDeclStmtValue, HirWhileStmt,
},
symbol::{FunctionId, FunctionSig, SymbolId, SymbolKind, SymbolTable},
types::SemaType,
},
};
pub struct Analyzer {
symbols: SymbolTable,
function_map: BTreeMap<String, FunctionId>,
functions: Vec<FunctionSig>,
current_func_return_type: Option<SemaType>,
diagnostic: Diagnositics,
loop_depth: usize,
}
impl Analyzer {
pub fn new() -> Self {
let mut analyzer = Self {
symbols: SymbolTable::new(),
function_map: BTreeMap::new(),
functions: vec![],
current_func_return_type: None,
diagnostic: Diagnositics::new(),
loop_depth: 0,
};
analyzer.declare_builtin_func("putint", vec![SemaType::I32], SemaType::Void);
analyzer.declare_builtin_func("getint", vec![], SemaType::I32);
analyzer
}
pub fn analyze(&mut self, compile_unit: CompileUnit) -> HirCompileUnit {
self.analyze_compile_unit(compile_unit)
}
pub fn get_diagnostics(&self) -> &Diagnositics {
&self.diagnostic
}
pub fn get_symbol_type(&self, symbol: SymbolId) -> SemaType {
self.symbols.get_symbol(symbol).ty
}
pub fn get_symbol_kind(&self, symbol: SymbolId) -> SymbolKind {
self.symbols.get_symbol(symbol).kind
}
pub fn get_function_sig(&self, function: FunctionId) -> &FunctionSig {
&self.functions[function.0]
}
fn declare_builtin_func(&mut self, name: &str, parameter_types: Vec<SemaType>, return_type: SemaType) {
let id = FunctionId(self.functions.len());
let sig = FunctionSig {
id,
name: name.to_string(),
return_type,
parameter_types,
};
self.function_map.insert(name.to_string(), id);
self.functions.push(sig);
}
fn add_error(&mut self, error: SemaError, span: Span) {
self.diagnostic.add_from_error(error, span);
}
fn analyze_compile_unit(&mut self, compile_unit: CompileUnit) -> HirCompileUnit {
let mut global_decls = vec![];
for decl in compile_unit.global_decls {
match decl {
GlobalDeclStmt::VarDecl(var_decl) => {
global_decls.push(HirGlobalDeclStmt::VarDecl(self.analyze_var_decl(var_decl, SymbolKind::Global)));
}
GlobalDeclStmt::FuncDecl(func_decl) => {
if let Some(func_decl) = self.analyze_func_decl(func_decl) {
global_decls.push(HirGlobalDeclStmt::FuncDecl(func_decl));
}
}
}
}
HirCompileUnit { global_decls }
}
fn analyze_var_decl(&mut self, var_decl: VarDeclStmt, kind: SymbolKind) -> HirVarDeclStmt {
let data_type = var_decl.data_type.into();
let mut values = vec![];
for value in var_decl.values {
match self.symbols.declare_variable(&value.name, kind, data_type) {
Ok(symbol) => values.push(HirVarDeclStmtValue {
symbol,
name_span: value.name_span,
}),
Err(e) => self.add_error(e, value.name_span),
}
}
HirVarDeclStmt {
values,
data_type,
type_span: var_decl.type_span,
}
}
fn analyze_func_decl(&mut self, func_decl: FuncDeclStmt) -> Option<HirFuncDeclStmt> {
if self.function_map.contains_key(&func_decl.name) {
self.add_error(SemaError::FunctionHasBeenDefined(func_decl.name.clone()), func_decl.name_span);
return None;
}
let function_id = FunctionId(self.functions.len());
let parameter_types = func_decl.params.iter().map(|param| param.param_type.into()).collect::<Vec<_>>();
let return_type = func_decl.return_type.into();
let sig = FunctionSig {
id: function_id,
name: func_decl.name.clone(),
return_type,
parameter_types,
};
self.function_map.insert(func_decl.name.clone(), function_id);
self.functions.push(sig.clone());
self.current_func_return_type = Some(return_type);
self.symbols.enter_scope();
let mut params = vec![];
for param in func_decl.params {
let param_type = param.param_type.into();
match self.symbols.declare_variable(&param.name, SymbolKind::Param, param_type) {
Ok(symbol) => params.push(HirParam {
symbol,
param_type,
name_span: param.name_span,
type_span: param.type_span,
}),
Err(e) => self.add_error(e, param.name_span),
}
}
let body = self.analyze_block_stmt(func_decl.body);
self.symbols.exit_scope();
self.current_func_return_type = None;
Some(HirFuncDeclStmt {
sig,
params,
body,
ret_type_span: func_decl.ret_type_span,
name_span: func_decl.name_span,
})
}
fn analyze_block_stmt(&mut self, block_stmt: BlockStmt) -> HirBlockStmt {
let mut statements = vec![];
for stmt in block_stmt.statements {
if let Some(stmt) = self.analyze_statement(stmt) {
statements.push(stmt);
}
}
HirBlockStmt { statements }
}
fn analyze_statement(&mut self, stmt: Statement) -> Option<HirStatement> {
match stmt {
Statement::Return(stmt) => Some(HirStatement::Return(self.analyze_return_stmt(stmt))),
Statement::If(stmt) => Some(HirStatement::If(self.analyze_if_stmt(stmt))),
Statement::While(stmt) => Some(HirStatement::While(self.analyze_while_stmt(stmt))),
Statement::Break(stmt) => Some(HirStatement::Break(self.analyze_break_stmt(stmt))),
Statement::Continue(stmt) => Some(HirStatement::Continue(self.analyze_continue_stmt(stmt))),
Statement::Block(stmt) => {
self.symbols.enter_scope();
let block = self.analyze_block_stmt(stmt);
self.symbols.exit_scope();
Some(HirStatement::Block(block))
}
Statement::Expr(expr) => self.analyze_expr(expr).map(HirStatement::Expr),
Statement::VarDecl(var_decl) => Some(HirStatement::VarDecl(self.analyze_var_decl(var_decl, SymbolKind::Local))),
}
}
fn analyze_return_stmt(&mut self, return_stmt: ReturnStmt) -> HirReturnStmt {
let expected_ty = self.current_func_return_type.unwrap();
let value = match return_stmt.value {
Some(expr) => {
if expected_ty == SemaType::Void {
self.add_error(SemaError::ReturnExpressionOnVoidFunction, return_stmt.span);
None
} else {
match self.analyze_expr(expr) {
Some(expr) => {
if expr.ty == SemaType::Void {
self.add_error(SemaError::InvalidOperand(SemaType::Void), return_stmt.span);
} else if expr.ty != expected_ty {
self.add_error(SemaError::TypeMismatch(expected_ty, expr.ty), return_stmt.span);
}
Some(expr)
}
None => None,
}
}
}
None => {
if expected_ty != SemaType::Void {
self.add_error(SemaError::TypeMismatch(expected_ty, SemaType::Void), return_stmt.span);
}
None
}
};
HirReturnStmt {
value,
span: return_stmt.span,
}
}
fn analyze_if_stmt(&mut self, if_stmt: IfStmt) -> HirIfStmt {
let condition = self.analyze_condition_expr(if_stmt.condition);
let then_branch = self.analyze_block_stmt(if_stmt.then_branch);
let mut ifelse_branch = vec![];
for branch in if_stmt.ifelse_branch {
let IfElseBranch { condition, then_branch } = branch;
ifelse_branch.push(HirIfElseBranch {
condition: self.analyze_condition_expr(condition),
then_branch: self.analyze_block_stmt(then_branch),
});
}
let else_branch = if_stmt.else_branch.map(|block| self.analyze_block_stmt(block));
HirIfStmt {
condition,
then_branch,
ifelse_branch,
else_branch,
}
}
fn analyze_while_stmt(&mut self, while_stmt: WhileStmt) -> HirWhileStmt {
let condition = self.analyze_condition_expr(while_stmt.condition);
self.loop_depth += 1;
let body = self.analyze_block_stmt(while_stmt.body);
self.loop_depth -= 1;
HirWhileStmt {
condition,
body,
}
}
fn analyze_break_stmt(&mut self, stmt: BreakStmt) -> HirBreakStmt {
if self.loop_depth == 0 {
self.add_error(SemaError::BreakOutsideLoop, stmt.span);
}
HirBreakStmt { span: stmt.span }
}
fn analyze_continue_stmt(&mut self, stmt: ContinueStmt) -> HirContinueStmt {
if self.loop_depth == 0 {
self.add_error(SemaError::ContinueOutsideLoop, stmt.span);
}
HirContinueStmt { span: stmt.span }
}
fn analyze_condition_expr(&mut self, expr: Expr) -> HirExpr {
let span = expr.span;
match self.analyze_expr(expr) {
Some(expr) => {
if expr.ty == SemaType::Void {
self.add_error(SemaError::InvalidOperand(SemaType::Void), span);
}
expr
}
None => HirExpr {
value: HirExprValue::IntLit(0),
ty: SemaType::I32,
span,
},
}
}
fn analyze_expr(&mut self, expr: Expr) -> Option<HirExpr> {
let span = expr.span;
match expr.value {
ExprValue::IntLit(value) => Some(HirExpr {
value: HirExprValue::IntLit(value),
ty: SemaType::I32,
span,
}),
ExprValue::Var(name) => {
let symbol = match self.symbols.get_variable(&name) {
Some(symbol) => symbol,
None => {
self.add_error(SemaError::VariableNotFound(name), span);
return None;
}
};
Some(HirExpr {
value: HirExprValue::Var(symbol),
ty: self.symbols.get_symbol(symbol).ty,
span,
})
}
ExprValue::Assign { lvalue, rvalue } => self.analyze_assign_expr(*lvalue, *rvalue, span),
ExprValue::UnaryOp { op, operand } => {
let operand_span = operand.span;
let operand = self.analyze_expr(*operand)?;
if operand.ty == SemaType::Void {
self.add_error(SemaError::InvalidOperand(SemaType::Void), operand_span);
return None;
}
let ty = match op {
crate::ast::types::UnaryOp::Not => SemaType::I1,
_ => operand.ty,
};
Some(HirExpr {
value: HirExprValue::UnaryOp {
op,
operand: Box::new(operand),
},
ty,
span,
})
}
ExprValue::BinaryOp { lhs, op, rhs } => self.analyze_binary_expr(*lhs, op, *rhs, span),
ExprValue::FuncCall(func_name, args) => self.analyze_func_call_expr(func_name, args, span),
}
}
fn analyze_assign_expr(&mut self, lvalue: Expr, rvalue: Expr, span: Span) -> Option<HirExpr> {
if !matches!(lvalue.value, ExprValue::Var(_)) {
self.add_error(SemaError::InvalidAssignmentTarget, lvalue.span);
return None;
}
let lvalue = self.analyze_expr(lvalue)?;
let rvalue_span = rvalue.span;
let rvalue = self.analyze_expr(rvalue)?;
if rvalue.ty == SemaType::Void {
self.add_error(SemaError::InvalidOperand(SemaType::Void), rvalue_span);
return None;
}
Some(HirExpr {
ty: lvalue.ty,
value: HirExprValue::Assign {
lvalue: Box::new(lvalue),
rvalue: Box::new(rvalue),
},
span,
})
}
fn analyze_binary_expr(&mut self, lhs: Expr, op: BinaryOp, rhs: Expr, span: Span) -> Option<HirExpr> {
let lhs_span = lhs.span;
let rhs_span = rhs.span;
let lhs = self.analyze_expr(lhs)?;
let rhs = self.analyze_expr(rhs)?;
if lhs.ty == SemaType::Void {
self.add_error(SemaError::InvalidOperand(SemaType::Void), lhs_span);
return None;
}
if rhs.ty == SemaType::Void {
self.add_error(SemaError::InvalidOperand(SemaType::Void), rhs_span);
return None;
}
let result_ty = if op.is_logical() {
SemaType::I1
} else {
match SemaType::get_elevate_result(lhs.ty, rhs.ty) {
Some(_) if op.is_cmp() => SemaType::I1,
Some(ty) => ty,
None => {
self.add_error(SemaError::IncompatiableOperand(lhs.ty, rhs.ty), lhs_span);
self.add_error(SemaError::IncompatiableOperand(lhs.ty, rhs.ty), rhs_span);
return None;
}
}
};
Some(HirExpr {
value: HirExprValue::BinaryOp {
lhs: Box::new(lhs),
op,
rhs: Box::new(rhs),
},
ty: result_ty,
span,
})
}
fn analyze_func_call_expr(&mut self, func_name: String, args: Vec<Expr>, span: Span) -> Option<HirExpr> {
let func_id = match self.function_map.get(&func_name).cloned() {
Some(func_id) => func_id,
None => {
self.add_error(SemaError::FunctionNotFound(func_name), span);
return None;
}
};
let func_def = self.functions[func_id.0].clone();
if args.len() < func_def.parameter_types.len() {
self.add_error(SemaError::TooFewArguments(func_def.parameter_types.len(), args.len()), span);
return None;
}
if args.len() > func_def.parameter_types.len() {
self.add_error(SemaError::TooManyArguments(func_def.parameter_types.len(), args.len()), span);
return None;
}
let mut has_error = false;
for parameter_type in &func_def.parameter_types {
if matches!(parameter_type, SemaType::Void) {
self.add_error(SemaError::InvalidParameterType(SemaType::Void), span);
has_error = true;
}
}
if has_error {
return None;
}
let mut hir_args = vec![];
for (i, arg) in args.into_iter().enumerate() {
let arg = self.analyze_expr(arg)?;
let parameter_type = func_def.parameter_types[i];
if parameter_type != arg.ty {
self.add_error(SemaError::TypeMismatch(parameter_type, arg.ty), span);
has_error = true;
continue;
}
hir_args.push(arg);
}
if has_error {
return None;
}
Some(HirExpr {
value: HirExprValue::FuncCall(func_id, hir_args),
ty: func_def.return_type,
span,
})
}
}
+35
View File
@@ -0,0 +1,35 @@
use thiserror::Error;
use crate::sema::types::SemaType;
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum SemaError {
#[error("variable `{0}` not found")]
VariableNotFound(String),
#[error("variable `{0}` has already been defined")]
VariableHasBeenDefined(String),
#[error("function `{0}` not found")]
FunctionNotFound(String),
#[error("function `{0}` has already been defined")]
FunctionHasBeenDefined(String),
#[error("incompatible operands: {0} and {1}")]
IncompatiableOperand(SemaType, SemaType),
#[error("invalid operand type: {0}")]
InvalidOperand(SemaType),
#[error("too few arguments: expected {0}, got {1}")]
TooFewArguments(usize, usize),
#[error("too many arguments: expected {0}, got {1}")]
TooManyArguments(usize, usize),
#[error("type mismatch, expected {0}, got {1}")]
TypeMismatch(SemaType, SemaType),
#[error("invalid assignment target")]
InvalidAssignmentTarget,
#[error("break statement outside of loop")]
BreakOutsideLoop,
#[error("continue statement outside of loop")]
ContinueOutsideLoop,
#[error("invalid parameter type: {0}")]
InvalidParameterType(SemaType),
#[error("return expression on void function")]
ReturnExpressionOnVoidFunction,
}
+106
View File
@@ -0,0 +1,106 @@
use crate::{ast::types::{BinaryOp, UnaryOp}, diagnostic::span::Span, sema::{symbol::{FunctionId, FunctionSig, SymbolId}, types::SemaType}};
pub struct HirCompileUnit {
pub global_decls: Vec<HirGlobalDeclStmt>,
}
pub enum HirGlobalDeclStmt {
VarDecl(HirVarDeclStmt),
FuncDecl(HirFuncDeclStmt),
}
pub struct HirVarDeclStmt {
pub values: Vec<HirVarDeclStmtValue>,
pub data_type: SemaType,
pub type_span: Span,
}
pub struct HirVarDeclStmtValue {
pub symbol: SymbolId,
pub name_span: Span,
}
pub struct HirFuncDeclStmt {
pub sig: FunctionSig,
pub params: Vec<HirParam>,
pub body: HirBlockStmt,
pub ret_type_span: Span,
pub name_span: Span,
}
pub struct HirParam {
pub symbol: SymbolId,
pub param_type: SemaType,
pub name_span: Span,
pub type_span: Span,
}
pub struct HirBlockStmt {
pub statements: Vec<HirStatement>,
}
pub enum HirStatement {
Return(HirReturnStmt),
Block(HirBlockStmt),
Expr(HirExpr),
VarDecl(HirVarDeclStmt),
If(HirIfStmt),
While(HirWhileStmt),
Break(HirBreakStmt),
Continue(HirContinueStmt),
}
pub struct HirIfStmt {
pub condition: HirExpr,
pub then_branch: HirBlockStmt,
pub ifelse_branch: Vec<HirIfElseBranch>,
pub else_branch: Option<HirBlockStmt>,
}
pub struct HirIfElseBranch {
pub condition: HirExpr,
pub then_branch: HirBlockStmt,
}
pub struct HirWhileStmt {
pub condition: HirExpr,
pub body: HirBlockStmt,
}
pub struct HirBreakStmt {
pub span: Span,
}
pub struct HirContinueStmt {
pub span: Span,
}
pub struct HirReturnStmt {
pub value: Option<HirExpr>,
pub span: Span,
}
pub struct HirExpr {
pub value: HirExprValue,
pub ty: SemaType,
pub span: Span,
}
pub enum HirExprValue {
IntLit(i64),
Var(SymbolId),
BinaryOp {
lhs: Box<HirExpr>,
op: BinaryOp,
rhs: Box<HirExpr>,
},
UnaryOp {
op: UnaryOp,
operand: Box<HirExpr>,
},
FuncCall(FunctionId, Vec<HirExpr>),
Assign {
lvalue: Box<HirExpr>,
rvalue: Box<HirExpr>,
},
}
+5
View File
@@ -0,0 +1,5 @@
pub mod analyzer;
pub mod err;
pub mod hir;
pub mod symbol;
pub mod types;
+83
View File
@@ -0,0 +1,83 @@
use std::collections::{BTreeMap, BTreeSet};
use crate::sema::{err::SemaError, types::SemaType};
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct SymbolId(pub usize);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SymbolKind {
Global,
Local,
Param,
}
#[derive(Clone, Debug)]
pub struct Symbol {
pub id: SymbolId,
pub name: String,
pub ty: SemaType,
pub kind: SymbolKind,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct FunctionId(pub usize);
#[derive(Clone, Debug)]
pub struct FunctionSig {
pub id: FunctionId,
pub name: String,
pub return_type: SemaType,
pub parameter_types: Vec<SemaType>,
}
pub struct SymbolTable {
symbols: Vec<Symbol>,
variable_map: BTreeMap<String, Vec<SymbolId>>,
scopes: Vec<BTreeSet<String>>,
}
impl SymbolTable {
pub fn new() -> Self {
Self {
symbols: vec![],
variable_map: BTreeMap::new(),
scopes: vec![BTreeSet::new()],
}
}
pub fn enter_scope(&mut self) {
self.scopes.push(BTreeSet::new());
}
pub fn exit_scope(&mut self) {
let variables = self.scopes.pop().unwrap();
for var in variables {
self.variable_map.get_mut(&var).unwrap().pop();
}
}
pub fn declare_variable(&mut self, name: &str, kind: SymbolKind, ty: SemaType) -> Result<SymbolId, SemaError> {
if self.scopes.last().unwrap().contains(name) {
return Err(SemaError::VariableHasBeenDefined(name.to_string()));
}
let id = SymbolId(self.symbols.len());
self.symbols.push(Symbol {
id,
name: name.to_string(),
ty,
kind,
});
self.variable_map.entry(name.to_string()).or_default().push(id);
self.scopes.last_mut().unwrap().insert(name.to_string());
Ok(id)
}
pub fn get_variable(&self, name: &str) -> Option<SymbolId> {
self.variable_map.get(name).and_then(|vars| vars.last()).cloned()
}
pub fn get_symbol(&self, id: SymbolId) -> &Symbol {
&self.symbols[id.0]
}
}
+51
View File
@@ -0,0 +1,51 @@
use std::fmt::Display;
use crate::{ast::types::Type as AstType, ir::types::IRType};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SemaType {
I32,
I1,
Void,
}
impl SemaType {
pub fn get_elevate_result(lhs: SemaType, rhs: SemaType) -> Option<SemaType> {
if lhs == rhs {
Some(lhs)
} else if (lhs == SemaType::I32 && rhs == SemaType::I1) || (lhs == SemaType::I1 && rhs == SemaType::I32) {
Some(SemaType::I32)
} else {
None
}
}
}
impl Display for SemaType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SemaType::I32 => write!(f, "i32"),
SemaType::I1 => write!(f, "i1"),
SemaType::Void => write!(f, "void"),
}
}
}
impl From<AstType> for SemaType {
fn from(value: AstType) -> Self {
match value {
AstType::Int => SemaType::I32,
AstType::Void => SemaType::Void,
}
}
}
impl From<SemaType> for IRType {
fn from(value: SemaType) -> Self {
match value {
SemaType::I32 => IRType::I32,
SemaType::I1 => IRType::I1,
SemaType::Void => IRType::Void,
}
}
}