Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
146 changes: 113 additions & 33 deletions zjit/src/backend/lir.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
use std::cell::Cell;
use std::collections::hash_map::Entry;
use std::collections::{HashMap, HashSet, VecDeque};
use std::collections::{HashMap, VecDeque};
use std::fmt;
use std::mem::take;
use std::rc::Rc;
Expand Down Expand Up @@ -2629,6 +2629,61 @@ impl Assembler
/// that area matters for two reasons: stack maps locate pushed survivors
/// at fixed frame offsets (see [StackState::stack_idx_for_caller_saved_reg]),
/// and the pops after the call must read back the exact slots the pushes wrote.
fn caller_saved_reg_survivors(
intervals: &[Interval],
call_sites: &[(usize, Option<VRegId>, Vec<VRegId>)],
alloc_regs: &RegPool,
) -> Vec<Vec<VRegId>> {
let mut calls_by_position: Vec<(usize, usize)> = call_sites.iter()
.enumerate()
.map(|(call_idx, (position, _, _))| (*position, call_idx))
.collect();
calls_by_position.sort_unstable_by_key(|(position, _)| *position);

let mut survivors: Vec<Vec<VRegId>> = (0..call_sites.len())
.map(|_| Vec::new())
.collect();

for interval in intervals {
let is_register = interval.assigned.get()
.and_then(|alloc| alloc.alloc_pool_index(alloc_regs))
.is_some();
if !is_register {
continue;
}

for range in &interval.ranges {
let start = calls_by_position.partition_point(|(position, _)| *position < range.from);
let end = calls_by_position.partition_point(|(position, _)| *position < range.to);
for &(_, call_idx) in &calls_by_position[start..end] {
if call_sites[call_idx].1 != Some(interval.vreg_id) {
survivors[call_idx].push(interval.vreg_id);
}
}
}
}

// Stack-map operands must be preserved even when they are not live
// across the call according to their interval.
for (call_idx, (_, _, stack_map_vregs)) in call_sites.iter().enumerate() {
for &vreg_id in stack_map_vregs {
if intervals[vreg_id].assigned.get()
.and_then(|alloc| alloc.alloc_pool_index(alloc_regs))
.is_some()
{
if !survivors[call_idx].contains(&vreg_id) {
survivors[call_idx].push(vreg_id);
}
}
}
}

survivors.into_iter().map(|mut survivors| {
survivors.sort_unstable();
survivors
}).collect()
}

pub fn handle_caller_saved_regs(
&mut self,
intervals: &[Interval],
Expand All @@ -2638,7 +2693,28 @@ impl Assembler
use crate::backend::parcopy;
use crate::backend::current::{C_RET_OPND, SCRATCH_REG, NATIVE_STACK_PTR};

for block_id in self.block_order() {
let block_order = self.block_order();
let mut call_sites = Vec::new();
for block_id in &block_order {
let block = &self.basic_blocks[block_id.0];
for (insn, insn_id) in block.insns.iter().zip(&block.insn_ids) {
if let Insn::CCall { data } = insn {
let position = insn_id.map(|id| id.0).unwrap_or(0);
let out_vreg_id = data.out.is_vreg().then(|| data.out.vreg_idx());
let stack_map_vregs = data.stack_map.as_ref().map_or_else(Vec::new, |stack_map| {
stack_map.stack.iter().filter_map(|entry| match entry {
StackMapEntry::Opnd(Opnd::VReg { idx, .. }) => Some(*idx),
_ => None,
}).collect()
});
call_sites.push((position, out_vreg_id, stack_map_vregs));
}
}
}
let caller_saved_survivors = Self::caller_saved_reg_survivors(intervals, &call_sites, alloc_regs);
let mut caller_saved_survivors = caller_saved_survivors.into_iter();

for block_id in block_order {
let block = &mut self.basic_blocks[block_id.0];
let old_insns = take(&mut block.insns);
let old_ids = take(&mut block.insn_ids);
Expand All @@ -2650,48 +2726,19 @@ impl Assembler
if let Insn::CCall { data } = insn {
let CCallData { opnds, stack_map, out, start_marker, end_marker, fptr } = *data;
let insn_number = insn_id.map(|id| id.0).unwrap_or(0);
let survivors = caller_saved_survivors.next()
.expect("CCall survivor list should be available");
// Do we have a case where a ccall is emitted, but nobody
// uses the result?
let call_result_live = out.is_vreg()
&& intervals[out.vreg_idx()].has_bounds()
&& !intervals[out.vreg_idx()].is_dead();

// Build a set of VRegIds that can be referenced by JITFrame for materializing the VM stack
let stack_vreg_ids: HashSet<VRegId> = if let Some(StackMap { stack, .. }) = &stack_map {
stack.iter().filter_map(|entry| match entry {
StackMapEntry::Opnd(Opnd::VReg { idx, .. }) => Some(*idx),
_ => None,
}).collect()
} else {
HashSet::default()
};

// Find survivors: intervals that are live across this Call
// instruction. We need to preserve the "surviving" registers
// past the ccall, so we're going to push them all on the
// stack, then pop after we make the ccall
let out_vreg_id = out.is_vreg().then(|| out.vreg_idx());
debug_assert!(
out_vreg_id.is_none_or(|id| !intervals[id].has_bounds() || intervals[id].born_at(insn_number)),
"a CCall's output interval must start at the CCall"
);
let survivors: Vec<VRegId> = intervals.iter()
.filter(|interval| {
// We need to spill register intervals on this CCall in two cases:
// 1) The VReg is live across the CCall. The VReg this CCall
// defines is not one of them: its range starts here, so it
// holds no value yet and there is nothing to preserve.
let live_across_call = Some(interval.vreg_id) != out_vreg_id
&& interval.covers(insn_number);

// 2) The VReg is referenced by the stack map for the CCall
let stack_map_reg = stack_vreg_ids.contains(&interval.vreg_id);
let is_register = interval.assigned.get().and_then(|alloc| alloc.alloc_pool_index(alloc_regs)).is_some();
is_register && (live_across_call || stack_map_reg)
})
.map(|interval| interval.vreg_id)
.collect();

let survivor_regs: Vec<Opnd> = survivors.iter()
.map(|&s| Opnd::Reg(intervals[s].assigned.get().unwrap().assigned_reg(alloc_regs).unwrap()))
.collect();
Expand Down Expand Up @@ -4385,6 +4432,39 @@ mod tests {
RegPool::with_allocatable(regs, num_allocatable)
}

#[test]
fn test_caller_saved_reg_survivors() {
let mut intervals: Vec<Interval> = (0..4).map(|idx| Interval::new(idx.into())).collect();
intervals[0].add_range(10, 20);
intervals[0].assigned.set(Some(Allocation::Reg(0)));
intervals[1].add_range(35, 40);
intervals[1].assigned.set(Some(Allocation::Reg(1)));
intervals[2].add_range(12, 20);
intervals[2].add_range(30, 35);
intervals[2].assigned.set(Some(Allocation::Reg(1)));
intervals[3].add_range(10, 40);
intervals[3].assigned.set(Some(Allocation::Stack(0)));

let call_sites = vec![
(35, Some(VRegId(1)), vec![]),
(25, None, vec![VRegId(0), VRegId(0), VRegId(3)]),
(15, None, vec![]),
(32, None, vec![]),
];
let survivors = Assembler::caller_saved_reg_survivors(
&intervals,
&call_sites,
&alloc_reg_pool(2),
);

assert_eq!(survivors, vec![
vec![],
vec![VRegId(0)],
vec![VRegId(0), VRegId(2)],
vec![VRegId(2)],
]);
}

#[test]
fn test_size_of_insn() {
// PatchPoint (PatchPointData), CCall (CCallData), and Target (Block/SideExit
Expand Down
Loading