alduin 0.0.1

WIP: A toy compiler backend
Documentation
use std::marker::PhantomData;

use bitvec::{bitvec, vec::BitVec};

use crate::compiler::graph::cfg::{
    liveness::{LiveInterval, LiveIntervalValueCategory},
    CFG,
};

use super::{Reg, ISA};

/// See: https://link.springer.com/content/pdf/10.1007/3-540-45937-5_17.pdf
pub struct LinearScanRegisterAllocator<'cfg, Isa: ISA> {
    cfg: &'cfg mut CFG<Isa::Op>,
    unhandled: Vec<usize>,
    active: Vec<usize>,
    inactive: Vec<usize>,
    handled: Vec<usize>,
    free: BitVec,
    weights: Vec<usize>,
    _p: PhantomData<Isa>,
}

impl<'cfg, Isa: ISA> LinearScanRegisterAllocator<'cfg, Isa> {
    pub fn new(cfg: &'cfg mut CFG<Isa::Op>) -> Self {
        cfg.used_registers = bitvec![0; Isa::Reg::MAX_COUNT];
        Self {
            cfg,
            unhandled: vec![],
            active: vec![],
            inactive: vec![],
            handled: vec![],
            free: bitvec![0; Isa::Reg::MAX_COUNT],
            weights: vec![0; Isa::Reg::MAX_COUNT],
            _p: PhantomData,
        }
    }

    fn interval(&self, index: usize) -> &LiveInterval {
        &self.cfg.liveness[index]
    }

    fn initialize(&mut self) {
        let mut unhandled = vec![];
        for (i, interval) in self.cfg.liveness.iter().enumerate() {
            assert!(!interval.ranges.is_empty());
            unhandled.push(i);
        }
        unhandled.sort_by_key(|i| self.interval(*i).ranges[0].start);
        self.unhandled = unhandled;
        self.active = vec![];
        self.inactive = vec![];
        self.handled = vec![];
        for gpr in Isa::Reg::GPRS {
            self.free.set((*gpr).into(), true);
        }
        for gpr in Isa::Reg::FPRS {
            self.free.set((*gpr).into(), true);
        }
        self.assign_stack_location_for_all_intervals();
    }

    fn assign_stack_location_for_all_intervals(&mut self) {
        let mut offset = 0;
        for interval in &mut *self.cfg.liveness {
            interval.mem = offset as i32;
            offset += interval.max_mem_size;
        }
        let align_mask = (1 << 4) - 1;
        self.cfg.stack_size = (offset + align_mask) & !align_mask;
    }

    fn allocate_mem_loc(&mut self, curr_index: usize) {
        for i in 0..self.weights.len() {
            self.weights[i] = 0;
        }
        let mut f = |i: usize| {
            if self.cfg.liveness[i].intersects(&self.cfg.liveness[curr_index]) {
                let i_reg = self.cfg.liveness[i].reg.unwrap();
                self.weights[i_reg] += self.cfg.liveness[i].weight;
            }
        };
        for i in 0..self.active.len() {
            f(self.active[i]);
        }
        for i in 0..self.inactive.len() {
            f(self.inactive[i]);
        }
        for i in 0..self.unhandled.len() {
            let it = self.unhandled[i];
            if self.cfg.liveness[it].reg.is_some() {
                f(it);
            }
        }
        // find r with minimum weights[r]
        let mut min_index = 0;
        let mut min_value = usize::MAX;
        let v_cat = self.cfg.liveness[curr_index].value_category.unwrap();
        for i in 0..self.weights.len() {
            if v_cat == LiveIntervalValueCategory::Int
                && (Isa::Reg::from(i).is_fpr()
                    || Isa::Reg::RESERVED_GPRS.contains(&Isa::Reg::from(i)))
            {
                continue;
            }
            if v_cat == LiveIntervalValueCategory::Float
                && (Isa::Reg::from(i).is_gpr()
                    || Isa::Reg::RESERVED_FPRS.contains(&Isa::Reg::from(i)))
            {
                continue;
            }
            if self.weights[i] < min_value {
                min_index = i;
                min_value = self.weights[i];
            }
        }
        assert_ne!(min_value, usize::MAX);
        let selected_reg = min_index;
        if self.interval(curr_index).weight < self.weights[selected_reg]
            || self.interval(curr_index).reg.is_some()
        {
            // if let Some(selected_reg) = self.interval(curr_index).reg {
            //     // move all active or inactive intervals to which r was assigned to handled
            //     // and assign memory locations to them
            //     for index in std::mem::take(&mut self.active) {
            //         if self.interval(index).reg == Some(selected_reg) {
            //             self.assign_mem(index);
            //             self.handled.push(index);
            //         } else {
            //             self.active.push(index);
            //         }
            //     }
            //     for index in std::mem::take(&mut self.inactive) {
            //         if self.interval(index).reg == Some(selected_reg) {
            //             self.assign_mem(index);
            //             self.handled.push(index);
            //         } else {
            //             self.inactive.push(index);
            //         }
            //     }
            // }
            // assign a memory location to cur and move cur to handled
            self.assign_mem(curr_index);
            self.handled.push(curr_index);
        } else {
            // move all active or inactive intervals to which r was assigned to handled
            // and assign memory locations to them
            for index in std::mem::take(&mut self.active) {
                if self.interval(index).reg == Some(selected_reg) {
                    self.assign_mem(index);
                    self.handled.push(index);
                } else {
                    self.active.push(index);
                }
            }
            for index in std::mem::take(&mut self.inactive) {
                if self.interval(index).reg == Some(selected_reg) {
                    self.assign_mem(index);
                    self.handled.push(index);
                } else {
                    self.inactive.push(index);
                }
            }
            // assign selected_reg to curr
            self.assign_reg(curr_index, Some(selected_reg));
            // move curr to active
            self.active.push(curr_index);
        }
    }

    /// Assign a stack location to an interval.
    /// This also removes the previously assigned register.
    fn assign_mem(&mut self, interval_index: usize) {
        // All intervals are pre-assigned with a stack location.
        // We only need to un-assign the register here
        self.assign_reg(interval_index, None);
    }

    /// Assign a register to an interval
    fn assign_reg(&mut self, interval_index: usize, reg: Option<usize>) {
        if let Some(reg) = reg {
            self.cfg.used_registers.set(reg, true);
        }
        self.cfg.liveness[interval_index].reg = reg;
    }

    pub fn allocate(&mut self) {
        self.initialize();
        while !self.unhandled.is_empty() {
            let curr_index = self.unhandled.remove(0);
            let curr = self.interval(curr_index);
            // check for internval in `active` that are handled or inactive
            let curr_begin = curr.ranges[0].start;
            for index in std::mem::take(&mut self.active) {
                let it = self.interval(index);
                let it_reg = it.reg.clone();
                if it.ranges.last().unwrap().end <= curr_begin {
                    self.handled.push(index);
                    if let Some(reg) = it_reg {
                        self.free.set(reg, true);
                    }
                } else if !it.covers(curr_begin) {
                    self.inactive.push(index);
                    if let Some(reg) = it_reg {
                        self.free.set(reg, true);
                    }
                } else {
                    self.active.push(index);
                }
            }
            // check for intervals in `inactive` that are handled or active
            for index in std::mem::take(&mut self.inactive) {
                let it = self.interval(index);
                let it_reg = it.reg.clone();
                if it.ranges.last().unwrap().end <= curr_begin {
                    self.handled.push(index);
                } else if it.covers(curr_begin) {
                    self.active.push(index);
                    if let Some(reg) = it_reg {
                        self.free.set(reg, false);
                    }
                } else {
                    self.inactive.push(index);
                }
            }
            // collect available registers
            let curr = self.interval(curr_index);
            let mut free_regs = self.free.clone();
            for i in &self.inactive {
                let it = self.interval(*i);
                if it.intersects(curr) {
                    if let Some(reg) = it.reg {
                        free_regs.set(reg, false)
                    }
                }
            }
            for i in &self.unhandled {
                let it = self.interval(*i);
                if it.reg.is_some() && it.intersects(curr) {
                    free_regs.set(it.reg.unwrap(), false);
                }
            }
            if curr.value_category.unwrap() == LiveIntervalValueCategory::Float {
                // remove all GPRs
                for r in Isa::Reg::GPRS {
                    free_regs.set((*r).into(), false);
                }
            } else {
                // remove all FPRs
                for r in Isa::Reg::FPRS {
                    free_regs.set((*r).into(), false);
                }
            }
            // select a register
            if !free_regs.any() || curr.reg.map(|r| !free_regs[r]).unwrap_or(false) {
                self.allocate_mem_loc(curr_index);
            } else {
                let reg = if curr.reg.is_none() {
                    let r = free_regs.first_one().unwrap();
                    self.assign_reg(curr_index, Some(r));
                    r
                } else {
                    assert!(
                        free_regs[curr.reg.unwrap()],
                        "Ref for fixed interval #{} is not free",
                        curr_index
                    );
                    curr.reg.unwrap()
                };
                self.free.set(reg, false);
                self.active.push(curr_index);
            }
        }
        // self.cfg.liveness.dump(self.cfg);
    }
}