use rucc_base::float::Format;
use super::{Arg, Call, Kind, Pass, Piece, Shape, Slot};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Class {
None,
Integer,
Sse,
X87,
Memory,
}
fn merge(left: Class, right: Class) -> Class {
match (left, right) {
(a, b) if a == b => a,
(Class::None, other) | (other, Class::None) => other,
(Class::Memory, _) | (_, Class::Memory) => Class::Memory,
(Class::X87, _) | (_, Class::X87) => Class::Memory,
(Class::Integer, _) | (_, Class::Integer) => Class::Integer,
_ => Class::Sse,
}
}
fn classes(shape: &Shape<'_>) -> Option<Vec<Class>> {
if shape.size > 16 {
return None;
}
let mut classes = vec![Class::None; usize::try_from(shape.size.div_ceil(8)).ok()?];
for piece in shape.pieces {
if piece.scalar.align > 1 && piece.offset % piece.scalar.align != 0 {
return None;
}
let class = match piece.scalar.kind {
Kind::Integer => Class::Integer,
Kind::Float(Format::X87Extended) => Class::X87,
Kind::Float(_) => Class::Sse,
};
for at in piece.offset / 8..=(piece.end() - 1) / 8 {
let slot = classes.get_mut(usize::try_from(at).ok()?)?;
*slot = merge(*slot, class);
}
}
classes.iter().all(|class| *class != Class::Memory).then_some(classes)
}
fn cost(classes: &[Class]) -> (u32, u32) {
let count = |want: Class| classes.iter().filter(|class| **class == want).count() as u32;
(count(Class::Integer) + count(Class::None), count(Class::Sse))
}
fn slots(shape: &Shape<'_>, classes: &[Class]) -> Vec<Slot> {
classes
.iter()
.enumerate()
.map(|(index, class)| {
let offset = index as u64 * 8;
let bytes = (shape.size - offset).min(8);
match class {
Class::Sse if bytes <= 4 => Slot::Float { offset, format: Format::Single },
Class::Sse => Slot::Float { offset, format: Format::Double },
_ => Slot::Integer { offset, size: u32::try_from(bytes).unwrap_or(8) },
}
})
.collect()
}
fn all_x87(shape: &Shape<'_>) -> bool {
let x87 = |piece: &Piece| piece.scalar.kind == Kind::Float(Format::X87Extended);
!shape.pieces.is_empty() && shape.pieces.iter().all(x87)
}
pub(super) fn returns(call: &mut Call, arg: &Arg<'_>) -> Pass {
let shape = match arg {
Arg::Void => return Pass::Ignore,
Arg::Scalar(_) => return Pass::Direct,
Arg::Aggregate(shape) => shape,
};
if shape.size == 0 {
return Pass::Ignore;
}
if all_x87(shape) && (shape.pieces.len() == 1 || (shape.pieces.len() == 2 && shape.complex)) {
let stack = shape
.pieces
.iter()
.map(|piece| Slot::Float { offset: piece.offset, format: Format::X87Extended });
return Pass::Pieces(stack.collect());
}
let Some(classes) = classes(shape) else { return sret(call) };
if classes.contains(&Class::X87) {
return sret(call);
}
Pass::Pieces(slots(shape, &classes))
}
fn sret(call: &mut Call) -> Pass {
call.gp = call.gp.saturating_sub(1);
Pass::Reference
}
pub(super) fn argument(call: &mut Call, arg: &Arg<'_>) -> Pass {
let shape = match arg {
Arg::Void => return Pass::Ignore,
Arg::Scalar(scalar) => {
match scalar.kind {
Kind::Integer => call.gp = call.gp.saturating_sub(registers(scalar.size)),
Kind::Float(Format::X87Extended) => {}
Kind::Float(_) => call.fp = call.fp.saturating_sub(1),
}
return Pass::Direct;
}
Arg::Aggregate(shape) => shape,
};
if shape.size == 0 {
return Pass::Ignore;
}
let Some(classes) = classes(shape) else { return Pass::Memory };
if classes.contains(&Class::X87) {
return Pass::Memory;
}
let (gp, fp) = cost(&classes);
if gp > call.gp || fp > call.fp {
return Pass::Memory;
}
call.gp -= gp;
call.fp -= fp;
Pass::Pieces(slots(shape, &classes))
}
fn registers(size: u64) -> u32 {
u32::try_from(size.div_ceil(8)).unwrap_or(1).max(1)
}
#[cfg(test)]
mod tests {
use super::super::tests::{float, fpr, gpr, int, packed, record, target};
use super::super::{Arg, Kind, Pass, Piece, Scalar, Shape};
use super::*;
fn call() -> Call {
target("x86_64-unknown-linux-gnu").call()
}
#[test]
fn a_structure_of_two_integers_travels_in_two_registers() {
let pieces = packed(&[int(4), int(4)]);
let shape = record(&pieces);
assert_eq!(shape.size, 8);
assert_eq!(call().argument(&Arg::Aggregate(shape)), Pass::Pieces(vec![gpr(0, 8)]));
let pieces = packed(&[int(4), int(4), int(4)]);
let shape = record(&pieces);
assert_eq!(shape.size, 12);
assert_eq!(
call().argument(&Arg::Aggregate(shape)),
Pass::Pieces(vec![gpr(0, 8), gpr(8, 4)])
);
}
#[test]
fn an_integer_beside_a_float_sends_the_float_into_a_general_register() {
let pieces = packed(&[int(4), float(Format::Single, 4)]);
assert_eq!(
call().argument(&Arg::Aggregate(record(&pieces))),
Pass::Pieces(vec![gpr(0, 8)])
);
let pieces = packed(&[int(8), float(Format::Double, 8)]);
assert_eq!(
call().argument(&Arg::Aggregate(record(&pieces))),
Pass::Pieces(vec![gpr(0, 8), fpr(8, Format::Double)])
);
}
#[test]
fn two_floats_in_one_eightbyte_are_one_vector_register() {
let pieces = packed(&[float(Format::Single, 4), float(Format::Single, 4)]);
assert_eq!(
call().argument(&Arg::Aggregate(record(&pieces))),
Pass::Pieces(vec![fpr(0, Format::Double)])
);
let pieces = packed(&[float(Format::Single, 4)]);
assert_eq!(
call().argument(&Arg::Aggregate(record(&pieces))),
Pass::Pieces(vec![fpr(0, Format::Single)])
);
}
#[test]
fn anything_over_sixteen_bytes_is_passed_in_memory() {
let pieces = packed(&[int(8), int(8), int(8)]);
assert_eq!(call().argument(&Arg::Aggregate(record(&pieces))), Pass::Memory);
assert_eq!(call().returns(&Arg::Aggregate(record(&pieces))), Pass::Reference);
}
#[test]
fn a_member_that_is_not_aligned_puts_the_whole_thing_in_memory() {
let pieces = [
Piece { offset: 0, scalar: int(1) },
Piece { offset: 1, scalar: Scalar { kind: Kind::Integer, size: 4, align: 4 } },
];
let shape = Shape { size: 5, align: 1, pieces: &pieces, complex: false };
assert_eq!(call().argument(&Arg::Aggregate(shape)), Pass::Memory);
let pieces = [
Piece { offset: 0, scalar: int(1) },
Piece { offset: 1, scalar: Scalar { kind: Kind::Integer, size: 4, align: 1 } },
];
let shape = Shape { size: 5, align: 1, pieces: &pieces, complex: false };
assert_eq!(call().argument(&Arg::Aggregate(shape)), Pass::Pieces(vec![gpr(0, 5)]));
}
#[test]
fn a_long_double_in_a_record_is_memory_going_in_and_the_x87_stack_coming_back() {
let pieces = packed(&[float(Format::X87Extended, 16)]);
let shape = record(&pieces);
assert_eq!(shape.size, 16);
assert_eq!(call().argument(&Arg::Aggregate(shape)), Pass::Memory);
assert_eq!(
call().returns(&Arg::Aggregate(shape)),
Pass::Pieces(vec![fpr(0, Format::X87Extended)])
);
}
#[test]
fn a_complex_long_double_comes_back_where_a_record_of_two_does_not() {
let pieces = packed(&[float(Format::X87Extended, 16), float(Format::X87Extended, 16)]);
let complex = Shape { complex: true, ..record(&pieces) };
assert_eq!(
call().returns(&Arg::Aggregate(complex)),
Pass::Pieces(vec![fpr(0, Format::X87Extended), fpr(16, Format::X87Extended)])
);
assert_eq!(call().returns(&Arg::Aggregate(record(&pieces))), Pass::Reference);
}
#[test]
fn an_aggregate_that_runs_out_of_registers_goes_to_memory_and_a_scalar_does_not() {
let pieces = packed(&[int(8), int(8)]);
let shape = Arg::Aggregate(record(&pieces));
let mut call = call();
for _ in 0..5 {
assert_eq!(call.argument(&Arg::Scalar(int(4))), Pass::Direct);
}
assert_eq!(call.argument(&shape), Pass::Memory);
assert_eq!(call.argument(&Arg::Scalar(int(4))), Pass::Direct);
assert_eq!(call.argument(&Arg::Scalar(int(4))), Pass::Direct);
}
#[test]
fn a_returned_pointer_to_memory_spends_the_register_the_first_argument_wanted() {
let big = packed(&[int(8), int(8), int(8)]);
let pair = packed(&[int(8), int(8)]);
let mut call = call();
assert_eq!(call.returns(&Arg::Aggregate(record(&big))), Pass::Reference);
for _ in 0..3 {
assert_eq!(call.argument(&Arg::Scalar(int(4))), Pass::Direct);
}
assert_eq!(
call.argument(&Arg::Aggregate(record(&pair))),
Pass::Pieces(vec![gpr(0, 8), gpr(8, 8)])
);
}
#[test]
fn an_int128_takes_two_registers_and_a_long_double_takes_none() {
let pair = packed(&[int(8), int(8)]);
let mut call = call();
for _ in 0..2 {
assert_eq!(call.argument(&Arg::Scalar(int(16))), Pass::Direct);
}
assert_eq!(call.argument(&Arg::Scalar(float(Format::X87Extended, 16))), Pass::Direct);
assert_eq!(
call.argument(&Arg::Aggregate(record(&pair))),
Pass::Pieces(vec![gpr(0, 8), gpr(8, 8)])
);
}
#[test]
fn an_aggregate_of_no_size_travels_nowhere() {
let shape = Shape { size: 0, align: 1, pieces: &[], complex: false };
assert_eq!(call().argument(&Arg::Aggregate(shape)), Pass::Ignore);
assert_eq!(call().returns(&Arg::Aggregate(shape)), Pass::Ignore);
assert_eq!(call().returns(&Arg::Void), Pass::Ignore);
}
}