use std::{
borrow::Cow,
collections::{HashMap, HashSet},
convert::Infallible,
error::Error,
fmt::{Display, Formatter}
};
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use crate::{
CanAllocate as _, Optimizer as _, Parser, Passes, SourceSpan,
StandardOptimizer, Validator,
ast::{self, ASTVisitor, Binding, Constant, Event, Expression, Node, Walk},
ir::{
AddressingMode, Immediate, Instruction, RegisterIndex,
RollingRecordIndex
},
parser::ParseError
};
pub fn compile_unoptimized(
source: &str
) -> Result<Function, CompilationError<'_>>
{
let ast = Parser::parse(source).map_err(CompilationError::ParseError)?;
Validator::validate(&ast)?;
Ok(Compiler::compile(&ast))
}
#[cfg_attr(doc, aquamarine::aquamarine)]
pub fn compile(source: &str) -> Result<Function, CompilationError<'_>>
{
let ast = Parser::parse(source).map_err(CompilationError::ParseError)?;
Validator::validate(&ast)?;
let function = Compiler::compile(&ast);
let optimizer = StandardOptimizer::new(Passes::all());
let function = optimizer
.optimize(function)
.map_err(|_| CompilationError::OptimizationFailed)?;
Ok(function)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CompilationError<'src>
{
ParseError(ParseError<'src>),
DuplicateParameter
{
name: Cow<'src, str>,
first: SourceSpan,
duplicate: SourceSpan
},
BindingCollidesWithParameter
{
name: Cow<'src, str>,
parameter: SourceSpan,
binding: SourceSpan
},
DuplicateBinding
{
name: Cow<'src, str>,
first: SourceSpan,
duplicate: SourceSpan
},
UseBeforeBind
{
name: Cow<'src, str>,
reference: SourceSpan,
binding: SourceSpan
},
OptimizationFailed
}
impl Display for CompilationError<'_>
{
fn fmt(&self, f: &mut Formatter) -> std::fmt::Result
{
match self
{
CompilationError::ParseError(e) =>
{
write!(f, "{}", e)
},
CompilationError::DuplicateParameter {
name,
first,
duplicate
} =>
{
write!(
f,
"duplicate parameter '{}' at {} (first declared at {})",
name, duplicate, first
)
},
CompilationError::BindingCollidesWithParameter {
name,
parameter,
binding
} =>
{
write!(
f,
"local binding '{}' at {} collides with formal \
parameter declared at {}",
name, binding, parameter
)
},
CompilationError::DuplicateBinding {
name,
first,
duplicate
} =>
{
write!(
f,
"duplicate local binding '{}' at {} (first bound at {})",
name, duplicate, first
)
},
CompilationError::UseBeforeBind {
name,
reference,
binding
} =>
{
write!(
f,
"reference to '{}' at {} precedes its binding at {}",
name, reference, binding
)
},
CompilationError::OptimizationFailed =>
{
write!(f, "optimization failed")
}
}
}
}
impl Error for CompilationError<'_> {}
pub struct Compiler<'a>
{
instructions: Vec<Instruction>,
next_register: RegisterIndex,
next_rolling_record: RollingRecordIndex,
arity: usize,
variables: HashMap<&'a str, RegisterIndex>,
bindings: HashMap<&'a str, AddressingMode>
}
impl<'a> Compiler<'a>
{
pub fn compile(ast: &'a ast::Function<'_>) -> Function
{
let mut compiler = Self {
instructions: Vec::new(),
next_register: RegisterIndex(0),
next_rolling_record: RollingRecordIndex(0),
arity: 0,
variables: HashMap::new(),
bindings: HashMap::new()
};
let _ = ast.accept(&mut compiler);
compiler.finish()
}
fn finish(self) -> Function
{
let mut parameters = Vec::new();
let mut externals = Vec::new();
for (name, register) in &self.variables
{
match register.0 >= self.arity
{
false => parameters.push((name, register)),
true => externals.push((name, register))
}
}
parameters.sort_by_key(|(_, register)| register.0);
externals.sort_by_key(|(_, register)| register.0);
let parameters = parameters
.into_iter()
.map(|(name, _)| name.to_string())
.collect();
let externals = externals
.into_iter()
.map(|(name, _)| name.to_string())
.collect();
Function {
parameters,
externals,
register_count: self.next_register.0,
rolling_record_count: self.next_rolling_record.0,
instructions: self.instructions
}
}
fn variable(&mut self, name: &'a str) -> RegisterIndex
{
match self.variables.get(name)
{
Some(®ister) => register,
None =>
{
let register = self.allocate_register();
self.variables.insert(name, register);
register
}
}
}
#[inline]
fn allocate_register(&mut self) -> RegisterIndex
{
self.next_register.allocate()
}
#[inline]
fn allocate_rolling_record(&mut self) -> RollingRecordIndex
{
self.next_rolling_record.allocate()
}
#[inline]
fn emit(&mut self, instruction: Instruction)
{
self.instructions.push(instruction);
}
fn generate_binary(
&mut self,
op1: AddressingMode,
op2: AddressingMode,
constructor: fn(
RegisterIndex,
AddressingMode,
AddressingMode
) -> Instruction
) -> AddressingMode
{
let dest = self.allocate_register();
self.emit(constructor(dest, op1, op2));
dest.into()
}
}
impl<'a, 'src: 'a> ASTVisitor<'a, 'src> for Compiler<'a>
{
type Output = AddressingMode;
type Error = Infallible;
fn enter_function(
&mut self,
node: &'a ast::Function<'src>
) -> Result<(), Infallible>
{
if let Some(ref parameters) = node.parameters
{
for param in parameters
{
self.variable(¶m.name);
}
self.arity = self.variables.len();
}
let binding_names = collect_binding_names(&node.body);
let externals = discover_externals(&node.body);
for external in externals
{
if !binding_names.contains(external)
{
self.variable(external);
}
}
Ok(())
}
fn visit_function(
&mut self,
_node: &'a ast::Function<'src>,
body: AddressingMode
) -> Result<AddressingMode, Infallible>
{
self.emit(Instruction::r#return(body));
Ok(body)
}
fn visit_group(
&mut self,
_node: &'a ast::Group<'src>,
expression: AddressingMode
) -> Result<AddressingMode, Infallible>
{
Ok(expression)
}
fn visit_constant(
&mut self,
node: &'a Constant
) -> Result<AddressingMode, Infallible>
{
Ok(Immediate(node.value).into())
}
fn visit_variable(
&mut self,
node: &'a ast::Variable<'src>
) -> Result<AddressingMode, Infallible>
{
if let Some(&addr) = self.bindings.get(&*node.name)
{
return Ok(addr);
}
let register = self.variable(&node.name);
Ok(register.into())
}
fn visit_binding(
&mut self,
node: &'a Binding<'src>,
expression: AddressingMode
) -> Result<AddressingMode, Infallible>
{
self.bindings.insert(&node.name, expression);
Ok(expression)
}
fn visit_range(
&mut self,
_node: &'a ast::Range<'src>,
start: AddressingMode,
end: AddressingMode
) -> Result<AddressingMode, Infallible>
{
let dest = self.allocate_rolling_record();
self.emit(Instruction::roll_range(dest, start, end));
let sum = self.allocate_register();
self.emit(Instruction::sum_rolling_record(sum, dest));
Ok(sum.into())
}
fn visit_standard_dice(
&mut self,
_node: &'a ast::StandardDice<'src>,
count: AddressingMode,
faces: AddressingMode
) -> Result<AddressingMode, Infallible>
{
let dest = self.allocate_rolling_record();
self.emit(Instruction::roll_standard_dice(dest, count, faces));
Ok(dest.into())
}
fn visit_custom_dice(
&mut self,
node: &'a ast::CustomDice<'src>,
count: AddressingMode
) -> Result<AddressingMode, Infallible>
{
let dest = self.allocate_rolling_record();
self.emit(Instruction::roll_custom_dice(
dest,
count,
node.faces.clone()
));
Ok(dest.into())
}
fn visit_drop_lowest(
&mut self,
_node: &'a ast::DropLowest<'src>,
dice: AddressingMode,
drop: Option<AddressingMode>
) -> Result<AddressingMode, Infallible>
{
let record: RollingRecordIndex = dice
.try_into()
.expect("dice visitor must return RollingRecord");
let count = drop.unwrap_or(Immediate(1).into());
self.emit(Instruction::drop_lowest(record, count));
Ok(record.into())
}
fn visit_drop_highest(
&mut self,
_node: &'a ast::DropHighest<'src>,
dice: AddressingMode,
drop: Option<AddressingMode>
) -> Result<AddressingMode, Infallible>
{
let record: RollingRecordIndex = dice
.try_into()
.expect("dice visitor must return RollingRecord");
let count = drop.unwrap_or(Immediate(1).into());
self.emit(Instruction::drop_highest(record, count));
Ok(record.into())
}
fn visit_add(
&mut self,
_node: &'a ast::Add<'src>,
left: AddressingMode,
right: AddressingMode
) -> Result<AddressingMode, Infallible>
{
Ok(self.generate_binary(left, right, Instruction::add))
}
fn visit_sub(
&mut self,
_node: &'a ast::Sub<'src>,
left: AddressingMode,
right: AddressingMode
) -> Result<AddressingMode, Infallible>
{
Ok(self.generate_binary(left, right, Instruction::sub))
}
fn visit_mul(
&mut self,
_node: &'a ast::Mul<'src>,
left: AddressingMode,
right: AddressingMode
) -> Result<AddressingMode, Infallible>
{
Ok(self.generate_binary(left, right, Instruction::mul))
}
fn visit_div(
&mut self,
_node: &'a ast::Div<'src>,
left: AddressingMode,
right: AddressingMode
) -> Result<AddressingMode, Infallible>
{
Ok(self.generate_binary(left, right, Instruction::div))
}
fn visit_mod(
&mut self,
_node: &'a ast::Mod<'src>,
left: AddressingMode,
right: AddressingMode
) -> Result<AddressingMode, Infallible>
{
Ok(self.generate_binary(left, right, Instruction::r#mod))
}
fn visit_exp(
&mut self,
_node: &'a ast::Exp<'src>,
left: AddressingMode,
right: AddressingMode
) -> Result<AddressingMode, Infallible>
{
Ok(self.generate_binary(left, right, Instruction::exp))
}
fn visit_neg(
&mut self,
node: &'a ast::Neg<'src>,
operand: AddressingMode
) -> Result<AddressingMode, Infallible>
{
if let Expression::Constant(Constant { value, .. }) =
node.operand.as_ref()
{
return Ok(Immediate(value.saturating_neg()).into());
}
let dest = self.allocate_register();
self.emit(Instruction::neg(dest, operand));
Ok(dest.into())
}
fn visit_expression(
&mut self,
_node: &'a Expression<'src>,
output: AddressingMode
) -> Result<AddressingMode, Infallible>
{
match output
{
AddressingMode::RollingRecord(record) =>
{
let sum = self.allocate_register();
self.emit(Instruction::sum_rolling_record(sum, record));
Ok(sum.into())
},
other => Ok(other)
}
}
}
fn discover_externals<'a>(expr: &'a Expression<'_>) -> Vec<&'a str>
{
let mut externals: Vec<&'a str> = Vec::new();
for event in Walk::new(Node::Expression(expr))
{
if let Event::Enter(Node::Expression(Expression::Variable(v))) = event
{
externals.push(&v.name);
}
}
externals
}
fn collect_binding_names<'a>(expr: &'a Expression<'_>) -> HashSet<&'a str>
{
let mut names: HashSet<&'a str> = HashSet::new();
for event in Walk::new(Node::Expression(expr))
{
if let Event::Enter(Node::Expression(Expression::Binding(b))) = event
{
names.insert(&b.name);
}
}
names
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct Function
{
pub parameters: Vec<String>,
pub externals: Vec<String>,
pub register_count: usize,
pub rolling_record_count: usize,
pub instructions: Vec<Instruction>
}
impl Function
{
#[inline]
pub fn arity(&self) -> usize { self.parameters.len() }
}
impl Display for Function
{
fn fmt(&self, f: &mut Formatter) -> std::fmt::Result
{
write!(f, "Function(")?;
for (i, parameter) in self.parameters.iter().enumerate()
{
if i != 0
{
write!(f, ", ")?;
}
write!(f, "{{{}}}@{}", parameter, i)?;
}
writeln!(
f,
") r#{} ⚅#{}",
self.register_count, self.rolling_record_count
)?;
write!(f, "\textern[")?;
for (i, external) in self.externals.iter().enumerate()
{
if i != 0
{
write!(f, ", ")?;
}
write!(f, "{{{}}}@{}", external, i + self.parameters.len())?;
}
writeln!(f, "]")?;
writeln!(f, "\tbody:")?;
for instruction in &self.instructions
{
writeln!(f, "\t\t{}", instruction)?;
}
Ok(())
}
}