diff --git a/zjit/src/backend/lir.rs b/zjit/src/backend/lir.rs index d7ba3d5333265f..b80b270d2a66d5 100644 --- a/zjit/src/backend/lir.rs +++ b/zjit/src/backend/lir.rs @@ -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; @@ -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, Vec)], + alloc_regs: &RegPool, + ) -> Vec> { + 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> = (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], @@ -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); @@ -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 = 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 = 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 = survivors.iter() .map(|&s| Opnd::Reg(intervals[s].assigned.get().unwrap().assigned_reg(alloc_regs).unwrap())) .collect(); @@ -4385,6 +4432,39 @@ mod tests { RegPool::with_allocatable(regs, num_allocatable) } + #[test] + fn test_caller_saved_reg_survivors() { + let mut intervals: Vec = (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