use std::collections::{HashMap, HashSet};
use crate::{
Add, AddressingMode, CanAllocate as _, CanVisitInstructions as _, Div,
DropHighest, DropLowest, Exp, Function, Immediate, Instruction,
InstructionVisitor, Max, Mod, Mul, Neg, ProgramCounter, RegisterIndex,
Return, RollCustomDice, RollRange, RollStandardDice, RollingRecordIndex,
Sub, SumRollingRecord
};
use crate::Optimizer;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct StrengthReducer
{
replacements: HashMap<AddressingMode, AddressingMode>,
next_register: RegisterIndex,
next_rolling_record: RollingRecordIndex,
dropped: HashSet<RollingRecordIndex>,
faces: HashMap<RollingRecordIndex, i32>,
instructions: Vec<Instruction>
}
impl Optimizer<()> for StrengthReducer
{
fn optimize(mut self, mut function: Function) -> Result<Function, ()>
{
let start_register =
function.parameters.len() + function.externals.len();
self.next_register = RegisterIndex(start_register);
loop
{
self.dropped = function
.instructions
.iter()
.filter_map(|inst| match inst
{
Instruction::DropLowest(drop) => Some(drop.dest),
Instruction::DropHighest(drop) => Some(drop.dest),
_ => None
})
.collect();
for instruction in &function.instructions
{
instruction.visit(&mut self).unwrap();
}
if function.instructions == self.instructions
{
return Ok(function)
}
function.register_count = self.next_register.0;
function.rolling_record_count = self.next_rolling_record.0;
function.instructions = self.instructions.clone();
self.replacements.clear();
self.faces.clear();
self.next_register = RegisterIndex(start_register);
self.next_rolling_record = RollingRecordIndex::default();
self.instructions.clear();
}
}
}
impl InstructionVisitor<()> for StrengthReducer
{
fn visit_roll_range(&mut self, range: &RollRange) -> Result<(), ()>
{
let (start, end) =
(self.replacement(range.start), self.replacement(range.end));
if start == end && !self.dropped.contains(&range.dest)
{
self.replace(range.dest, start);
return Ok(());
}
let dest = self.next_rolling_record();
self.replace(range.dest, dest);
self.emit(RollRange { dest, start, end });
Ok(())
}
fn visit_roll_standard_dice(
&mut self,
roll: &RollStandardDice
) -> Result<(), ()>
{
let (count, faces) =
(self.replacement(roll.count), self.replacement(roll.faces));
if let AddressingMode::Immediate(Immediate(1)) = faces
{
self.reduce_to_count(roll.dest, roll.count, 1);
return Ok(());
}
if let AddressingMode::Immediate(Immediate(1)) = count
{
let dest = self.next_rolling_record();
self.replace(roll.dest, dest);
self.emit(RollRange {
dest,
start: Immediate(1).into(),
end: faces
});
return Ok(());
}
let dest = self.next_rolling_record();
self.replace(roll.dest, dest);
self.emit(RollStandardDice { dest, count, faces });
Ok(())
}
fn visit_roll_custom_dice(
&mut self,
roll: &RollCustomDice
) -> Result<(), ()>
{
if roll.distinct_faces() == 1
{
self.reduce_to_count(roll.dest, roll.count, roll.faces[0]);
return Ok(());
}
let mut faces = roll.faces.to_vec();
faces.sort();
let mut counter = faces[0];
let mut contiguous = true;
for face in &faces[1..]
{
if *face == counter + 1
{
counter = *face;
}
else
{
contiguous = false;
break;
}
}
if let AddressingMode::Immediate(Immediate(1)) =
self.replacement(roll.count)
&& contiguous
{
let dest = self.next_rolling_record();
self.replace(roll.dest, dest);
self.emit(RollRange {
dest,
start: Immediate(faces[0]).into(),
end: Immediate(faces[faces.len() - 1]).into()
});
return Ok(());
}
if contiguous && faces[0] == 1
{
let dest = self.next_rolling_record();
self.replace(roll.dest, dest);
self.emit(RollStandardDice {
dest,
count: self.replacement(roll.count),
faces: Immediate(faces[faces.len() - 1]).into()
});
return Ok(())
}
let dest = self.next_rolling_record();
self.replace(roll.dest, dest);
self.emit(RollCustomDice {
dest,
count: self.replacement(roll.count),
faces: roll.faces.clone()
});
Ok(())
}
fn visit_drop_lowest(&mut self, drop: &DropLowest) -> Result<(), ()>
{
match self.replacement(drop.dest)
{
AddressingMode::Register(_) | AddressingMode::Immediate(_) =>
{
self.drop_from_count(drop.dest, drop.count);
Ok(())
},
AddressingMode::RollingRecord(dest) =>
{
let count = self.replacement(drop.count);
if let AddressingMode::Immediate(Immediate(..=0)) = count
{
self.replace(drop.dest, dest);
return Ok(());
}
if let AddressingMode::Immediate(Immediate(count)) = count
&& let Some(pc) = self.find_drop_instruction(|inst| {
match DropLowest::try_from(inst.clone()).ok()
{
Some(drop) =>
{
drop.dest == dest
&& matches!(
drop.count,
AddressingMode::Immediate(_)
)
},
_ => false
}
})
{
let previous = self.instructions.remove(pc.0);
let AddressingMode::Immediate(Immediate(previous_count)) =
*previous.sources().last().unwrap()
else
{
unreachable!()
};
self.emit(DropLowest {
dest,
count: Immediate(previous_count.saturating_add(count))
.into()
});
return Ok(())
}
self.emit(DropLowest {
dest,
count: self.replacement(drop.count)
});
Ok(())
}
}
}
fn visit_drop_highest(&mut self, drop: &DropHighest) -> Result<(), ()>
{
match self.replacement(drop.dest)
{
AddressingMode::Register(_) | AddressingMode::Immediate(_) =>
{
self.drop_from_count(drop.dest, drop.count);
Ok(())
},
AddressingMode::RollingRecord(dest) =>
{
let count = self.replacement(drop.count);
if let AddressingMode::Immediate(Immediate(..=0)) = count
{
self.replace(drop.dest, dest);
return Ok(());
}
if let AddressingMode::Immediate(Immediate(count)) = count
&& let Some(pc) = self.find_drop_instruction(|inst| {
match DropHighest::try_from(inst.clone()).ok()
{
Some(drop) =>
{
drop.dest == dest
&& matches!(
drop.count,
AddressingMode::Immediate(_)
)
},
_ => false
}
})
{
let previous = self.instructions.remove(pc.0);
let AddressingMode::Immediate(Immediate(previous_count)) =
*previous.sources().last().unwrap()
else
{
unreachable!()
};
self.emit(DropHighest {
dest,
count: Immediate(previous_count.saturating_add(count))
.into()
});
return Ok(())
}
self.emit(DropHighest {
dest,
count: self.replacement(drop.count)
});
Ok(())
}
}
}
fn visit_sum_rolling_record(
&mut self,
sum: &SumRollingRecord
) -> Result<(), ()>
{
match self.replacement(sum.src)
{
src @ (AddressingMode::Immediate(_)
| AddressingMode::Register(_)) =>
{
let sum_value = match self.faces.get(&sum.src)
{
Some(&face) if face != 1 => match src
{
AddressingMode::Immediate(Immediate(count)) =>
{
Immediate(count.saturating_mul(face)).into()
},
count =>
{
let dest = self.next_register();
self.emit(Mul {
dest,
op1: count,
op2: Immediate(face).into()
});
dest.into()
}
},
_ => src
};
self.replace(sum.dest, sum_value);
Ok(())
},
AddressingMode::RollingRecord(_) =>
{
let dest = self.next_register();
self.replace(sum.dest, dest);
self.emit(SumRollingRecord {
dest,
src: self.replacement(sum.src).try_into().unwrap()
});
Ok(())
}
}
}
fn visit_add(&mut self, inst: &Add) -> Result<(), ()>
{
let (op1, op2) =
(self.replacement(inst.op1), self.replacement(inst.op2));
if let AddressingMode::Immediate(Immediate(0)) = op1
{
self.replace(inst.dest, op2);
return Ok(())
}
if let AddressingMode::Immediate(Immediate(0)) = op2
{
self.replace(inst.dest, op1);
return Ok(())
}
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Add { dest, op1, op2 });
Ok(())
}
fn visit_sub(&mut self, inst: &Sub) -> Result<(), ()>
{
let (op1, op2) =
(self.replacement(inst.op1), self.replacement(inst.op2));
if let AddressingMode::Immediate(Immediate(0)) = op1
{
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Neg { dest, op: op2 });
return Ok(())
}
if let AddressingMode::Immediate(Immediate(0)) = op2
{
self.replace(inst.dest, op1);
return Ok(())
}
if op1 == op2
{
self.replace(inst.dest, Immediate(0));
return Ok(())
}
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Sub { dest, op1, op2 });
Ok(())
}
fn visit_mul(&mut self, inst: &Mul) -> Result<(), ()>
{
let (op1, op2) =
(self.replacement(inst.op1), self.replacement(inst.op2));
if let AddressingMode::Immediate(Immediate(0)) = op1
{
self.replace(inst.dest, Immediate(0));
return Ok(())
}
if let AddressingMode::Immediate(Immediate(0)) = op2
{
self.replace(inst.dest, Immediate(0));
return Ok(())
}
if let AddressingMode::Immediate(Immediate(1)) = op1
{
self.replace(inst.dest, op2);
return Ok(())
}
if let AddressingMode::Immediate(Immediate(1)) = op2
{
self.replace(inst.dest, op1);
return Ok(())
}
if let AddressingMode::Immediate(Immediate(-1)) = op1
{
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Neg { dest, op: op2 });
return Ok(())
}
if let AddressingMode::Immediate(Immediate(-1)) = op2
{
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Neg { dest, op: op1 });
return Ok(())
}
if let AddressingMode::Immediate(Immediate(2)) = op1
{
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Add {
dest,
op1: op2,
op2
});
return Ok(())
}
if let AddressingMode::Immediate(Immediate(2)) = op2
{
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Add {
dest,
op1,
op2: op1
});
return Ok(())
}
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Mul { dest, op1, op2 });
Ok(())
}
fn visit_div(&mut self, inst: &Div) -> Result<(), ()>
{
let (op1, op2) =
(self.replacement(inst.op1), self.replacement(inst.op2));
if let AddressingMode::Immediate(Immediate(0)) = op1
{
self.replace(inst.dest, Immediate(0));
return Ok(())
}
if let AddressingMode::Immediate(Immediate(0)) = op2
{
self.replace(inst.dest, Immediate(0));
return Ok(())
}
if let AddressingMode::Immediate(Immediate(1)) = op2
{
self.replace(inst.dest, op1);
return Ok(())
}
if let AddressingMode::Immediate(Immediate(-1)) = op2
{
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Neg { dest, op: op1 });
return Ok(())
}
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Div { dest, op1, op2 });
Ok(())
}
fn visit_mod(&mut self, inst: &Mod) -> Result<(), ()>
{
let (op1, op2) =
(self.replacement(inst.op1), self.replacement(inst.op2));
if let AddressingMode::Immediate(Immediate(0)) = op1
{
self.replace(inst.dest, Immediate(0));
return Ok(())
}
if let AddressingMode::Immediate(Immediate(0)) = op2
{
self.replace(inst.dest, Immediate(0));
return Ok(())
}
if let AddressingMode::Immediate(Immediate(1)) = op2
{
self.replace(inst.dest, Immediate(0));
return Ok(())
}
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Mod { dest, op1, op2 });
Ok(())
}
fn visit_exp(&mut self, inst: &Exp) -> Result<(), ()>
{
let (op1, op2) =
(self.replacement(inst.op1), self.replacement(inst.op2));
if let AddressingMode::Immediate(Immediate(0)) = op2
{
self.replace(inst.dest, Immediate(1));
return Ok(())
}
if let AddressingMode::Immediate(Immediate(1)) = op2
{
self.replace(inst.dest, op1);
return Ok(())
}
if let AddressingMode::Immediate(Immediate(1)) = op1
{
self.replace(inst.dest, Immediate(1));
return Ok(())
}
if let AddressingMode::Immediate(Immediate(2)) = op2
{
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Mul {
dest,
op1,
op2: op1
});
return Ok(())
}
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Exp { dest, op1, op2 });
Ok(())
}
fn visit_max(&mut self, inst: &Max) -> Result<(), ()>
{
let op1 = self.replacement(inst.op1);
let op2 = self.replacement(inst.op2);
if let AddressingMode::Immediate(Immediate(i32::MIN)) = op1
{
self.replace(inst.dest, op2);
return Ok(())
}
if let AddressingMode::Immediate(Immediate(i32::MIN)) = op2
{
self.replace(inst.dest, op1);
return Ok(())
}
if let AddressingMode::Immediate(Immediate(i32::MAX)) = op1
{
self.replace(inst.dest, op1);
return Ok(())
}
if let AddressingMode::Immediate(Immediate(i32::MAX)) = op2
{
self.replace(inst.dest, op2);
return Ok(())
}
if op1 == op2
{
self.replace(inst.dest, op1);
return Ok(())
}
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Max { dest, op1, op2 });
Ok(())
}
fn visit_neg(&mut self, inst: &Neg) -> Result<(), ()>
{
let dest = self.next_register();
self.replace(inst.dest, dest);
self.emit(Neg {
dest,
op: self.replacement(inst.op)
});
Ok(())
}
fn visit_return(&mut self, inst: &Return) -> Result<(), ()>
{
self.emit(Return {
src: self.replacement(inst.src)
});
Ok(())
}
}
impl StrengthReducer
{
#[inline]
fn next_register(&mut self) -> RegisterIndex
{
self.next_register.allocate()
}
#[inline]
fn next_rolling_record(&mut self) -> RollingRecordIndex
{
self.next_rolling_record.allocate()
}
#[inline]
fn replace(
&mut self,
old: impl Into<AddressingMode>,
new: impl Into<AddressingMode>
)
{
self.replacements.insert(old.into(), new.into());
}
#[inline]
fn replacement(&self, op: impl Into<AddressingMode>) -> AddressingMode
{
let op: AddressingMode = op.into();
*self.replacements.get(&op).unwrap_or(&op)
}
fn find_drop_instruction(
&self,
filter: impl Fn(&Instruction) -> bool
) -> Option<ProgramCounter>
{
self.instructions.iter().enumerate().rev().find_map(
move |(pc, inst)| match filter(inst)
{
true => Some(pc.into()),
false => None
}
)
}
fn at_least_zero(&mut self, op: AddressingMode) -> AddressingMode
{
match op
{
AddressingMode::Immediate(Immediate(value)) =>
{
Immediate(value.max(0)).into()
},
op =>
{
let dest = self.next_register();
self.emit(Max {
dest,
op1: op,
op2: Immediate(0).into()
});
dest.into()
}
}
}
fn reduce_to_count(
&mut self,
dest: RollingRecordIndex,
count: AddressingMode,
face: i32
)
{
let count = self.at_least_zero(self.replacement(count));
self.replace(dest, count);
self.faces.insert(dest, face);
}
fn drop_from_count(
&mut self,
dest: RollingRecordIndex,
dropped: AddressingMode
)
{
debug_assert!(
self.faces.contains_key(&dest),
"a drop from a value that is not a count of dice"
);
let count = self.replacement(dest);
let dropped = self.at_least_zero(self.replacement(dropped));
let remaining = match (count, dropped)
{
(_, AddressingMode::Immediate(Immediate(0))) => count,
(
AddressingMode::Immediate(Immediate(count)),
AddressingMode::Immediate(Immediate(dropped))
) => Immediate(count.saturating_sub(dropped).max(0)).into(),
(count, dropped) =>
{
let difference = self.next_register();
self.emit(Sub {
dest: difference,
op1: count,
op2: dropped
});
self.at_least_zero(difference.into())
}
};
self.replace(dest, remaining);
}
#[inline]
fn emit(&mut self, inst: impl Into<Instruction>)
{
self.instructions.push(inst.into());
}
}