#[cfg(feature = "std")]
extern crate std;
use alloc::vec::Vec;
use core::cmp::Ordering;
const MIN_PARALLEL: usize = 65_536;
const MIN_CHUNK: usize = 16_384;
const MAX_PARTS: usize = 8;
static LEASED: core::sync::atomic::AtomicUsize = core::sync::atomic::AtomicUsize::new(0);
struct Lease(usize);
impl Drop for Lease {
fn drop(&mut self) {
LEASED.fetch_sub(self.0, core::sync::atomic::Ordering::Relaxed);
}
}
fn lease_amount(held: usize, want: usize, cap: usize) -> usize {
want.min(cap.saturating_sub(held))
}
fn take_lease(want: usize, cap: usize) -> Lease {
use core::sync::atomic::Ordering::Relaxed;
let mut held = LEASED.load(Relaxed);
loop {
let take = lease_amount(held, want, cap);
if take == 0 {
return Lease(0);
}
match LEASED.compare_exchange_weak(held, held + take, Relaxed, Relaxed) {
Ok(_) => return Lease(take),
Err(now) => held = now,
}
}
}
fn parts_for(n: usize, extra_workers: usize) -> usize {
if n < MIN_PARALLEL || extra_workers == 0 {
return 1;
}
(extra_workers + 1).min(n / MIN_CHUNK).min(MAX_PARTS)
}
#[cfg(not(feature = "std"))]
pub(crate) fn sort_total<T, F>(mut v: Vec<T>, _how: Workers, cmp: &F) -> Vec<T>
where
T: Send,
F: Fn(&T, &T) -> Ordering + Sync,
{
v.sort_unstable_by(|a, b| cmp(a, b));
v
}
#[cfg(not(feature = "std"))]
pub(crate) fn sort_total_stable<T, F>(mut v: Vec<T>, _how: Workers, cmp: &F) -> Vec<T>
where
T: Send,
F: Fn(&T, &T) -> Ordering + Sync,
{
v.sort_by(|a, b| cmp(a, b));
v
}
#[derive(Clone, Copy, Debug)]
pub(crate) struct Workers {
pub per_sort: usize,
pub per_process: usize,
}
impl Workers {
pub(crate) const fn serial() -> Self {
Self {
per_sort: 0,
per_process: 0,
}
}
}
#[cfg(feature = "std")]
pub(crate) fn sort_total<T, F>(v: Vec<T>, how: Workers, cmp: &F) -> Vec<T>
where
T: Send,
F: Fn(&T, &T) -> Ordering + Sync,
{
sort_in(v, how, cmp, false)
}
#[cfg(feature = "std")]
pub(crate) fn sort_total_stable<T, F>(v: Vec<T>, how: Workers, cmp: &F) -> Vec<T>
where
T: Send,
F: Fn(&T, &T) -> Ordering + Sync,
{
sort_in(v, how, cmp, true)
}
#[cfg(feature = "std")]
fn sort_in<T, F>(mut v: Vec<T>, how: Workers, cmp: &F, stable: bool) -> Vec<T>
where
T: Send,
F: Fn(&T, &T) -> Ordering + Sync,
{
let n = v.len();
let wanted = parts_for(n, how.per_sort);
let lease = take_lease(wanted.saturating_sub(1), how.per_process);
let parts = lease.0 + 1;
if parts < 2 {
if stable {
v.sort_by(|a, b| cmp(a, b));
} else {
v.sort_unstable_by(|a, b| cmp(a, b));
}
return v;
}
let base = n / parts;
let rem = n % parts;
let mut runs: Vec<Vec<T>> = Vec::with_capacity(parts);
let mut src = v.into_iter();
for i in 0..parts {
runs.push(src.by_ref().take(base + usize::from(i < rem)).collect());
}
std::thread::scope(|s| {
let mut rest: &mut [Vec<T>] = runs.as_mut_slice();
while let Some((head, tail)) = rest.split_first_mut() {
rest = tail;
s.spawn(move || {
if stable {
head.sort_by(|a, b| cmp(a, b));
} else {
head.sort_unstable_by(|a, b| cmp(a, b));
}
});
}
});
while runs.len() > 1 {
runs = merge_round(runs, cmp);
}
drop(lease);
runs.pop().unwrap_or_default()
}
#[cfg(feature = "std")]
fn merge_round<T, F>(runs: Vec<Vec<T>>, cmp: &F) -> Vec<Vec<T>>
where
T: Send,
F: Fn(&T, &T) -> Ordering + Sync,
{
let mut pairs: Vec<(Vec<T>, Option<Vec<T>>)> = Vec::with_capacity(runs.len().div_ceil(2));
let mut it = runs.into_iter();
while let Some(a) = it.next() {
pairs.push((a, it.next()));
}
std::thread::scope(|s| {
let handles: Vec<_> = pairs
.into_iter()
.map(|(a, b)| {
s.spawn(move || match b {
None => a,
Some(b) => merge_two(a, b, cmp),
})
})
.collect();
handles
.into_iter()
.map(|h| h.join().expect("a merge thread panicked"))
.collect()
})
}
#[cfg(feature = "std")]
fn merge_two<T, F>(a: Vec<T>, b: Vec<T>, cmp: &F) -> Vec<T>
where
F: Fn(&T, &T) -> Ordering,
{
enum Take {
A,
B,
RestA,
RestB,
}
let mut out = Vec::with_capacity(a.len() + b.len());
let mut ai = a.into_iter().peekable();
let mut bi = b.into_iter().peekable();
loop {
let take = match (ai.peek(), bi.peek()) {
(Some(x), Some(y)) => {
if cmp(x, y) == Ordering::Greater {
Take::B
} else {
Take::A
}
}
(Some(_), None) => Take::RestA,
(None, Some(_)) => Take::RestB,
(None, None) => return out,
};
match take {
Take::A => out.extend(ai.next()),
Take::B => out.extend(bi.next()),
Take::RestA => {
out.extend(ai);
return out;
}
Take::RestB => {
out.extend(bi);
return out;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
fn plenty(per_sort: usize) -> Workers {
Workers {
per_sort,
per_process: 64,
}
}
fn total(a: &(u64, u32), b: &(u64, u32)) -> Ordering {
a.0.cmp(&b.0).then_with(|| a.1.cmp(&b.1))
}
#[test]
fn every_fan_out_gives_the_serial_answer() {
for n in [0usize, 1, 2, 3, 65_535, 65_536, 70_001, 131_072] {
let v: Vec<(u64, u32)> = (0..n as u32)
.map(|i| ((u64::from(i) * 7919) % 1000, i))
.collect();
let mut want = v.clone();
want.sort_unstable_by(total);
for workers in [0usize, 1, 2, 3, 7, 64] {
let got = sort_total(v.clone(), plenty(workers), &total);
assert_eq!(got, want, "n={n} workers={workers}");
}
}
}
#[test]
fn an_odd_run_count_carries_the_last_one_through() {
let runs = vec![
vec![(1u64, 0u32), (3, 1)],
vec![(2u64, 2u32), (4, 3)],
vec![(0u64, 4u32)],
];
let out = merge_round(runs, &total);
assert_eq!(
out,
vec![vec![(1, 0), (2, 2), (3, 1), (4, 3)], vec![(0, 4)]]
);
}
#[test]
fn a_non_copy_element_sorts_the_same() {
let n = 70_000u32;
let v: Vec<(alloc::vec::Vec<u8>, u32)> = (0..n)
.map(|i| (((i * 7919) % 1000).to_be_bytes().to_vec(), i))
.collect();
let owned = |a: &(alloc::vec::Vec<u8>, u32), b: &(alloc::vec::Vec<u8>, u32)| {
a.0.cmp(&b.0).then_with(|| a.1.cmp(&b.1))
};
let mut want = v.clone();
want.sort_unstable_by(owned);
assert_eq!(sort_total(v, plenty(4), &owned), want);
}
#[test]
fn the_guc_decides_the_fan_out() {
assert_eq!(parts_for(1_000, 8), 1, "small inputs stay serial");
assert_eq!(parts_for(400_000, 0), 1, "the GUC can turn it off");
assert_eq!(
parts_for(400_000, 2),
3,
"PG's default of two workers sorts in three processes"
);
assert_eq!(parts_for(400_000, 3), 4);
assert_eq!(
parts_for(400_000, 64),
MAX_PARTS,
"a server runs more than one query"
);
assert_eq!(parts_for(70_000, 64), 4, "no chunk under MIN_CHUNK");
}
#[test]
fn the_process_cap_is_shared_between_sorts() {
assert_eq!(lease_amount(0, 6, 8), 6);
assert_eq!(lease_amount(6, 6, 8), 2, "what is left, not what was asked");
assert_eq!(
lease_amount(8, 6, 8),
0,
"a sort that finds none runs serially"
);
assert_eq!(lease_amount(9, 6, 8), 0, "and cannot go negative");
assert_eq!(lease_amount(0, 6, 0), 0, "a cap of zero lets nothing out");
}
#[test]
fn a_sort_with_no_threads_left_still_sorts() {
let v: Vec<(u64, u32)> = (0..70_000u32).map(|i| (u64::from(i % 997), i)).collect();
let mut want = v.clone();
want.sort_unstable_by(total);
assert_eq!(
sort_total(
v,
Workers {
per_sort: 8,
per_process: 0
},
&total
),
want
);
}
}