use std::{
borrow::Cow,
collections::{HashMap, HashSet},
convert::Infallible,
error::Error,
fmt::{Display, Formatter}
};
#[cfg(feature = "serde")]
use serde::{Deserialize, Deserializer, 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, is_canonical_name}
};
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))]
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() }
pub fn validate(&self) -> Result<(), FunctionError>
{
let mut first_by_name = HashMap::new();
for (index, name) in
self.parameters.iter().chain(&self.externals).enumerate()
{
if !is_canonical_name(name)
{
return Err(FunctionError::NonCanonicalName {
name: name.clone(),
index
})
}
if let Some(&first) = first_by_name.get(name.as_str())
{
return Err(FunctionError::DuplicateName {
name: name.clone(),
first,
index
})
}
first_by_name.insert(name.as_str(), index);
}
let declared_args = self.parameters.len() + self.externals.len();
if declared_args > self.register_count
{
return Err(FunctionError::InsufficientRegisterCount {
register_count: self.register_count,
required: declared_args
})
}
let mut register_seen = vec![false; self.register_count];
register_seen[..declared_args].fill(true);
let mut record_seen = vec![false; self.rolling_record_count];
for (instruction, inst) in self.instructions.iter().enumerate()
{
self.check_instruction(
inst,
instruction,
&mut register_seen,
&mut record_seen
)?;
}
if let Some(instruction) = self
.instructions
.iter()
.position(|inst| matches!(inst, Instruction::Return(_)))
&& instruction + 1 < self.instructions.len()
{
return Err(FunctionError::EarlyReturn { instruction })
}
if !matches!(self.instructions.last(), Some(Instruction::Return(_)))
{
return Err(FunctionError::MissingReturn)
}
if let Some(gap) = register_seen.iter().position(|seen| !*seen)
{
return Err(FunctionError::RegisterGap {
index: gap,
register_count: self.register_count
})
}
if let Some(gap) = record_seen.iter().position(|seen| !*seen)
{
return Err(FunctionError::RollingRecordGap {
index: gap,
rolling_record_count: self.rolling_record_count
})
}
Ok(())
}
fn check_instruction(
&self,
inst: &Instruction,
instruction: usize,
register_seen: &mut [bool],
record_seen: &mut [bool]
) -> Result<(), FunctionError>
{
let check_register = |idx: RegisterIndex,
seen: &mut [bool]|
-> Result<(), FunctionError> {
if idx.0 >= self.register_count
{
return Err(FunctionError::RegisterOutOfBounds {
index: idx.0,
register_count: self.register_count,
instruction
})
}
seen[idx.0] = true;
Ok(())
};
let check_record = |idx: RollingRecordIndex,
seen: &mut [bool]|
-> Result<(), FunctionError> {
if idx.0 >= self.rolling_record_count
{
return Err(FunctionError::RollingRecordOutOfBounds {
index: idx.0,
rolling_record_count: self.rolling_record_count,
instruction
})
}
seen[idx.0] = true;
Ok(())
};
let check_mode = |mode: AddressingMode,
seen: &mut [bool]|
-> Result<(), FunctionError> {
match mode
{
AddressingMode::Immediate(_) => Ok(()),
AddressingMode::Register(reg) => check_register(reg, seen),
AddressingMode::RollingRecord(_) =>
{
Err(FunctionError::UnexpectedRollingRecordOperand {
instruction
})
},
}
};
match inst
{
Instruction::RollRange(inst) =>
{
check_record(inst.dest, record_seen)?;
check_mode(inst.start, register_seen)?;
check_mode(inst.end, register_seen)?;
},
Instruction::RollStandardDice(inst) =>
{
check_record(inst.dest, record_seen)?;
check_mode(inst.count, register_seen)?;
check_mode(inst.faces, register_seen)?;
},
Instruction::RollCustomDice(inst) =>
{
check_record(inst.dest, record_seen)?;
check_mode(inst.count, register_seen)?;
if inst.faces.is_empty()
{
return Err(FunctionError::FacelessCustomDice {
instruction
})
}
},
Instruction::DropLowest(inst) =>
{
check_record(inst.dest, record_seen)?;
check_mode(inst.count, register_seen)?;
},
Instruction::DropHighest(inst) =>
{
check_record(inst.dest, record_seen)?;
check_mode(inst.count, register_seen)?;
},
Instruction::SumRollingRecord(inst) =>
{
check_register(inst.dest, register_seen)?;
check_record(inst.src, record_seen)?;
},
Instruction::Add(inst) =>
{
check_register(inst.dest, register_seen)?;
check_mode(inst.op1, register_seen)?;
check_mode(inst.op2, register_seen)?;
},
Instruction::Sub(inst) =>
{
check_register(inst.dest, register_seen)?;
check_mode(inst.op1, register_seen)?;
check_mode(inst.op2, register_seen)?;
},
Instruction::Mul(inst) =>
{
check_register(inst.dest, register_seen)?;
check_mode(inst.op1, register_seen)?;
check_mode(inst.op2, register_seen)?;
},
Instruction::Div(inst) =>
{
check_register(inst.dest, register_seen)?;
check_mode(inst.op1, register_seen)?;
check_mode(inst.op2, register_seen)?;
},
Instruction::Mod(inst) =>
{
check_register(inst.dest, register_seen)?;
check_mode(inst.op1, register_seen)?;
check_mode(inst.op2, register_seen)?;
},
Instruction::Exp(inst) =>
{
check_register(inst.dest, register_seen)?;
check_mode(inst.op1, register_seen)?;
check_mode(inst.op2, register_seen)?;
},
Instruction::Max(inst) =>
{
check_register(inst.dest, register_seen)?;
check_mode(inst.op1, register_seen)?;
check_mode(inst.op2, register_seen)?;
},
Instruction::Neg(inst) =>
{
check_register(inst.dest, register_seen)?;
check_mode(inst.op, register_seen)?;
},
Instruction::Return(inst) =>
{
check_mode(inst.src, register_seen)?;
}
}
Ok(())
}
}
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(())
}
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub enum FunctionError
{
NonCanonicalName
{
name: String,
index: usize
},
DuplicateName
{
name: String,
first: usize,
index: usize
},
InsufficientRegisterCount
{
register_count: usize,
required: usize
},
RegisterOutOfBounds
{
index: usize,
register_count: usize,
instruction: usize
},
RollingRecordOutOfBounds
{
index: usize,
rolling_record_count: usize,
instruction: usize
},
FacelessCustomDice
{
instruction: usize
},
UnexpectedRollingRecordOperand
{
instruction: usize
},
MissingReturn,
EarlyReturn
{
instruction: usize
},
RegisterGap
{
index: usize,
register_count: usize
},
RollingRecordGap
{
index: usize,
rolling_record_count: usize
}
}
impl FunctionError
{
pub fn instruction(&self) -> Option<usize>
{
match self
{
Self::RegisterOutOfBounds { instruction, .. }
| Self::RollingRecordOutOfBounds { instruction, .. }
| Self::FacelessCustomDice { instruction }
| Self::UnexpectedRollingRecordOperand { instruction }
| Self::EarlyReturn { instruction } => Some(*instruction),
Self::NonCanonicalName { .. }
| Self::DuplicateName { .. }
| Self::InsufficientRegisterCount { .. }
| Self::MissingReturn
| Self::RegisterGap { .. }
| Self::RollingRecordGap { .. } => None
}
}
pub(crate) fn describe(&self, f: &mut Formatter<'_>) -> std::fmt::Result
{
match self
{
Self::NonCanonicalName { name, index } => write!(
f,
"variable @{} is named `{}`, which is not canonical",
index, name
),
Self::DuplicateName { name, first, index } => write!(
f,
"variables @{} and @{} are both named `{}` (names must be \
distinct)",
first, index, name
),
Self::InsufficientRegisterCount {
register_count,
required
} => write!(
f,
"r#{} registers cannot hold the {} parameters and external \
variables",
register_count, required
),
Self::RegisterOutOfBounds {
index,
register_count,
..
} => write!(
f,
"register @{} exceeds declared register count r#{}",
index, register_count
),
Self::RollingRecordOutOfBounds {
index,
rolling_record_count,
..
} => write!(
f,
"rolling record ⚅{} exceeds declared rolling record count \
⚅#{}",
index, rolling_record_count
),
Self::FacelessCustomDice { .. } =>
{
write!(f, "custom dice must have at least one face")
},
Self::UnexpectedRollingRecordOperand { .. } =>
{
write!(f, "rolling record operand is not permitted here")
},
Self::MissingReturn =>
{
write!(f, "the function has no return")
},
Self::EarlyReturn { .. } => write!(
f,
"return is not the last instruction (a function ends with \
its only return)"
),
Self::RegisterGap {
index,
register_count,
..
} => write!(
f,
"register @{} is declared by r#{} but is never referenced \
(no gaps are permitted in the register file)",
index, register_count
),
Self::RollingRecordGap {
index,
rolling_record_count,
..
} => write!(
f,
"rolling record ⚅{} is declared by ⚅#{} but is never \
referenced (no gaps are permitted in the rolling record \
file)",
index, rolling_record_count
)
}
}
}
impl Display for FunctionError
{
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result
{
if let Some(instruction) = self.instruction()
{
write!(f, "instruction {}: ", instruction)?;
}
self.describe(f)
}
}
impl Error for FunctionError {}
#[cfg(feature = "serde")]
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct UncheckedFunction
{
parameters: Vec<String>,
externals: Vec<String>,
register_count: usize,
rolling_record_count: usize,
instructions: Vec<Instruction>
}
#[cfg(feature = "serde")]
impl<'de> Deserialize<'de> for Function
{
fn deserialize<D: Deserializer<'de>>(
deserializer: D
) -> Result<Self, D::Error>
{
use serde::de::Error as _;
let UncheckedFunction {
parameters,
externals,
register_count,
rolling_record_count,
instructions
} = UncheckedFunction::deserialize(deserializer)?;
let function = Self {
parameters,
externals,
register_count,
rolling_record_count,
instructions
};
function.validate().map_err(D::Error::custom)?;
Ok(function)
}
}