use std::collections::{BTreeMap, HashMap, HashSet, VecDeque};
use std::ffi::c_void;
use std::sync::{Condvar, Mutex};
use proofman_fields::PrimeField64;
use proofman_common::{MemoryHandlerRecursive, Proof, ProofType};
use crate::Ledger;
use proofman_starks_lib_c::{release_stream_reservation_c, reserve_best_stream_nonblock_c, reserve_stream_if_free_c};
type Key = (usize, usize, ProofType);
fn is_aggregation(t: ProofType) -> bool {
t == ProofType::Recursive1 || t == ProofType::Recursive2
}
const RECURSIVE_ORDER: [ProofType; 2] = [ProofType::Recursive2, ProofType::Recursive1];
pub struct StreamReservation {
d_buffers: usize,
stream_id: u32,
armed: bool,
}
unsafe impl Send for StreamReservation {}
impl StreamReservation {
fn new(d_buffers: usize, stream_id: u32) -> Self {
Self { d_buffers, stream_id, armed: true }
}
pub fn stream_id(&self) -> u32 {
self.stream_id
}
pub fn commit(mut self) {
self.armed = false;
}
}
impl Drop for StreamReservation {
fn drop(&mut self) {
if self.armed {
release_stream_reservation_c(self.d_buffers as *mut c_void, self.stream_id);
}
}
}
pub enum WorkerPick<F: PrimeField64> {
Recursive(Proof<F>, StreamReservation),
Basic(usize, StreamReservation),
}
fn candidate_keys(ready: &[(Key, usize)], order: &[ProofType]) -> Vec<Key> {
let mut out = Vec::new();
for &t in order {
let mut of_type: Vec<(Key, usize)> = ready.iter().copied().filter(|((_, _, kt), _)| *kt == t).collect();
of_type.sort_by(|(ka, ba), (kb, bb)| bb.cmp(ba).then((ka.0, ka.1).cmp(&(kb.0, kb.1))));
out.extend(of_type.into_iter().map(|(k, _)| k));
}
out
}
pub struct RecursiveScheduler<F: PrimeField64> {
d_buffers: usize,
queues: HashMap<Key, VecDeque<Proof<F>>>,
basic_queue: HashMap<(usize, usize), VecDeque<usize>>,
resident_keys: HashSet<Key>,
stream_warm: BTreeMap<usize, Key>,
}
impl<F: PrimeField64> RecursiveScheduler<F> {
pub fn new(d_buffers: *mut c_void) -> Self {
Self {
d_buffers: d_buffers as usize,
queues: HashMap::new(),
basic_queue: HashMap::new(),
resident_keys: HashSet::new(),
stream_warm: BTreeMap::new(),
}
}
fn d_buffers(&self) -> *mut c_void {
self.d_buffers as *mut c_void
}
pub fn push(&mut self, w: Proof<F>) {
let key = (w.airgroup_id, w.air_id, w.proof_type);
self.queues.entry(key).or_default().push_back(w);
}
pub fn is_empty(&self) -> bool {
self.queues.values().all(|q| q.is_empty())
}
fn pop_key(&mut self, key: Key) -> Option<Proof<F>> {
let q = self.queues.get_mut(&key)?;
let item = q.pop_front();
if q.is_empty() {
self.queues.remove(&key);
}
item
}
fn is_warm_somewhere(&self, key: Key) -> bool {
self.stream_warm.values().any(|k| *k == key)
}
fn next_of_types(&mut self, order: &[ProofType], force_recursive: bool) -> Option<(Proof<F>, StreamReservation)> {
let ready: Vec<(Key, usize)> =
self.queues.iter().filter(|(_, q)| !q.is_empty()).map(|(k, q)| (*k, q.len())).collect();
let candidates = candidate_keys(&ready, order);
if candidates.is_empty() {
return None;
}
let (key, s) = self.pick_and_reserve(&candidates, force_recursive)?;
let w = self.pop_key(key).expect("candidate key non-empty");
Some((w, s))
}
pub fn next_recursive(&mut self) -> Option<(Proof<F>, StreamReservation)> {
self.next_of_types(&RECURSIVE_ORDER, true)
}
pub fn next_nonrecursive(&mut self) -> Option<WorkerPick<F>> {
if let Some((w, s)) = self.next_of_types(&[ProofType::Compressor], false) {
return Some(WorkerPick::Recursive(w, s));
}
let allow_resident = self.is_empty();
if let Some((id, s)) = self.next_basic(allow_resident) {
return Some(WorkerPick::Basic(id, s));
}
if let Some((w, s)) = self.next_of_types(&RECURSIVE_ORDER, false) {
return Some(WorkerPick::Recursive(w, s));
}
None
}
pub fn next_basic(&mut self, include_resident: bool) -> Option<(usize, StreamReservation)> {
let mut ready: Vec<((usize, usize), usize)> = self
.basic_queue
.iter()
.filter(|(k, q)| {
!q.is_empty() && (include_resident || !self.resident_keys.contains(&(k.0, k.1, ProofType::Basic)))
})
.map(|(k, q)| (*k, q.len()))
.collect();
if ready.is_empty() {
return None;
}
ready.sort_by(|(ka, ba), (kb, bb)| bb.cmp(ba).then(ka.cmp(kb)));
let candidates: Vec<Key> = ready.iter().map(|((ag, air), _)| (*ag, *air, ProofType::Basic)).collect();
let (key, s) = self.pick_and_reserve(&candidates, false)?;
let (ag, air, _) = key;
let q = self.basic_queue.get_mut(&(ag, air)).expect("candidate basic key non-empty");
let id = q.pop_front().expect("candidate basic key non-empty");
if q.is_empty() {
self.basic_queue.remove(&(ag, air));
}
Some((id, s))
}
pub fn push_basic(&mut self, instance_id: usize, airgroup_id: usize, air_id: usize, resident: bool) {
if resident {
self.resident_keys.insert((airgroup_id, air_id, ProofType::Basic));
}
self.basic_queue.entry((airgroup_id, air_id)).or_default().push_back(instance_id);
}
pub fn basic_is_empty(&self) -> bool {
self.basic_queue.values().all(|q| q.is_empty())
}
fn pick_and_reserve(&mut self, candidates: &[Key], force_recursive: bool) -> Option<(Key, StreamReservation)> {
let d = self.d_buffers();
for &key in candidates {
let (ag, air, t) = key;
let type_str: &'static str = t.into();
for (&s, &k) in self.stream_warm.iter() {
if k == key && reserve_stream_if_free_c(d, s as u32, ag as u64, air as u64, type_str, force_recursive) {
return Some((key, StreamReservation::new(self.d_buffers, s as u32)));
}
}
}
for &key in candidates {
if self.is_warm_somewhere(key) {
continue; }
if let Some(s) = self.reserve_best(key, force_recursive) {
return Some((key, s));
}
}
for &key in candidates {
if let Some(s) = self.reserve_best(key, force_recursive) {
return Some((key, s));
}
}
None
}
fn reserve_best(&mut self, key: Key, force_recursive: bool) -> Option<StreamReservation> {
let (ag, air, t) = key;
let type_str: &'static str = t.into();
let s = reserve_best_stream_nonblock_c(
self.d_buffers(),
ag as u64,
air as u64,
type_str,
is_aggregation(t),
force_recursive,
);
if s == u32::MAX {
return None;
}
self.stream_warm.insert(s as usize, key);
Some(StreamReservation::new(self.d_buffers, s))
}
pub fn drain_all(&mut self) -> (Vec<Proof<F>>, Vec<usize>) {
let witnesses: Vec<Proof<F>> = self.queues.drain().flat_map(|(_, q)| q.into_iter()).collect();
let basics: Vec<usize> = self.basic_queue.drain().flat_map(|(_, q)| q.into_iter()).collect();
self.resident_keys.clear();
(witnesses, basics)
}
}
pub struct SharedScheduler<F: PrimeField64> {
pub lock: Mutex<RecursiveScheduler<F>>,
pub ready: Condvar,
}
impl<F: PrimeField64> SharedScheduler<F> {
pub fn new(inner: RecursiveScheduler<F>) -> Self {
Self { lock: Mutex::new(inner), ready: Condvar::new() }
}
pub fn push(&self, w: Proof<F>) {
self.lock.lock().unwrap().push(w);
self.ready.notify_all();
}
}
pub fn recover_drained_witnesses<F: PrimeField64 + Send + Sync + 'static>(
witnesses: Vec<Proof<F>>,
memory_handler_recursive_witness: &MemoryHandlerRecursive<F>,
ledger: &Ledger,
) -> usize {
let recovered = witnesses.len();
for mut w in witnesses {
let compressor = w.proof_type == ProofType::Compressor;
drop(memory_handler_recursive_witness.adopt_witness(std::mem::take(&mut w.circom_witness), compressor));
if let Some(idx) = w.global_idx {
ledger.settle(idx as u64, w.proof_type.as_usize());
}
}
recovered
}
pub fn eligible_stream_count(need: usize, class_sizes: &[usize]) -> usize {
class_sizes.iter().filter(|&&s| s >= need).count()
}
pub fn witness_slot_cap(eligible: usize, n_classes: usize) -> usize {
if eligible == 0 || n_classes == 0 || eligible >= n_classes {
return usize::MAX;
}
eligible + 1
}
pub fn schedule_key(
airgroup_id: usize,
air_id: usize,
is_stored: bool,
has_compressor: bool,
proof_cost: u64,
) -> (u8, std::cmp::Reverse<u64>, usize, usize) {
let priority_tier: u8 = if is_stored && has_compressor {
0
} else if is_stored {
1
} else if has_compressor {
2
} else {
3
};
(priority_tier, std::cmp::Reverse(proof_cost), airgroup_id, air_id)
}
#[cfg(test)]
mod tests {
use super::{candidate_keys, is_aggregation, Key, RECURSIVE_ORDER};
use proofman_common::ProofType;
const FEED_ORDER: [ProofType; 3] = [ProofType::Compressor, ProofType::Recursive2, ProofType::Recursive1];
fn k(ag: usize, air: usize, t: ProofType) -> Key {
(ag, air, t)
}
#[test]
fn aggregation_flag_matches_c() {
assert!(is_aggregation(ProofType::Recursive1));
assert!(is_aggregation(ProofType::Recursive2));
assert!(!is_aggregation(ProofType::Compressor));
}
#[test]
fn recursive_order_is_rec2_then_rec1() {
assert_eq!(RECURSIVE_ORDER, [ProofType::Recursive2, ProofType::Recursive1]);
}
#[test]
fn candidates_prefer_bigger_backlog_within_type() {
let ready = vec![(k(0, 5, ProofType::Recursive1), 2), (k(0, 2, ProofType::Recursive1), 9)];
let out = candidate_keys(&ready, &FEED_ORDER);
assert_eq!(out, vec![k(0, 2, ProofType::Recursive1), k(0, 5, ProofType::Recursive1)]);
}
#[test]
fn candidates_respect_type_priority() {
let ready = vec![
(k(0, 1, ProofType::Recursive1), 5),
(k(0, 2, ProofType::Compressor), 1),
(k(0, 3, ProofType::Recursive2), 1),
];
let out = candidate_keys(&ready, &FEED_ORDER);
assert_eq!(
out,
vec![k(0, 2, ProofType::Compressor), k(0, 3, ProofType::Recursive2), k(0, 1, ProofType::Recursive1),]
);
}
#[test]
fn candidates_recursive_order_drops_compressor() {
let ready = vec![(k(0, 2, ProofType::Compressor), 9), (k(0, 1, ProofType::Recursive1), 1)];
let out = candidate_keys(&ready, &RECURSIVE_ORDER);
assert_eq!(out, vec![k(0, 1, ProofType::Recursive1)]);
}
#[test]
fn candidates_tie_break_by_airgroup_air() {
let ready = vec![(k(0, 7, ProofType::Recursive1), 3), (k(0, 2, ProofType::Recursive1), 3)];
let out = candidate_keys(&ready, &FEED_ORDER);
assert_eq!(out, vec![k(0, 2, ProofType::Recursive1), k(0, 7, ProofType::Recursive1)]);
}
}
#[cfg(test)]
mod drain_tests {
use super::*;
use crate::{DeviceBuffersPtr, DeviceCompletions};
use proofman_fields::{Field, Goldilocks};
use proofman_common::MemoryHandlerRecursive;
type F = Goldilocks;
const W_SIZE: usize = 8;
const W_SIZE_COMPRESSOR: usize = 4;
fn scheduler() -> RecursiveScheduler<F> {
RecursiveScheduler::<F>::new(std::ptr::null_mut())
}
fn handler() -> MemoryHandlerRecursive<F> {
MemoryHandlerRecursive::new(2, 2, W_SIZE, W_SIZE_COMPRESSOR, W_SIZE, W_SIZE_COMPRESSOR)
}
fn witness(h: &MemoryHandlerRecursive<F>, t: ProofType, global_idx: usize) -> Proof<F> {
let buf = if t == ProofType::Compressor { h.take_buffer_witness_compressor() } else { h.take_buffer_witness() };
Proof::new_witness(t, 0, 0, Some(global_idx), buf, 1)
}
#[test]
fn drain_all_takes_everything_and_leaves_the_scheduler_empty() {
let h = handler();
let mut s = scheduler();
s.push(witness(&h, ProofType::Compressor, 0));
s.push(witness(&h, ProofType::Recursive1, 1));
s.push(witness(&h, ProofType::Recursive2, 2));
s.push_basic(10, 0, 0, false);
s.push_basic(11, 0, 1, true);
let (witnesses, basics) = s.drain_all();
assert_eq!(witnesses.len(), 3);
assert_eq!(basics.len(), 2);
assert!(s.is_empty() && s.basic_is_empty(), "drain must leave nothing behind");
let (w2, b2) = s.drain_all();
assert!(w2.is_empty() && b2.is_empty());
recover_drained_witnesses(witnesses, &h, &DeviceCompletions::new().acquire(null_ptr()).ledger());
}
fn null_ptr() -> DeviceBuffersPtr {
DeviceBuffersPtr(std::ptr::null_mut())
}
#[test]
fn recovery_returns_compressor_buffers_to_the_compressor_pool() {
let h = handler();
let mut s = scheduler();
s.push(witness(&h, ProofType::Compressor, 0));
s.push(witness(&h, ProofType::Compressor, 1));
let (witnesses, _) = s.drain_all();
let owner = DeviceCompletions::new().acquire(null_ptr());
assert_eq!(recover_drained_witnesses(witnesses, &h, &owner.ledger()), 2);
h.reset().expect("all four pools must be whole after recovery");
}
#[test]
fn recovery_settles_the_units_armed_at_hand_off() {
let h = handler();
let owner = DeviceCompletions::new().acquire(null_ptr());
let ledger = owner.ledger();
let mut s = scheduler();
for (t, idx) in [(ProofType::Compressor, 0usize), (ProofType::Recursive1, 1), (ProofType::Recursive2, 2)] {
ledger.arm(idx as u64, t.as_usize()).commit();
s.push(witness(&h, t, idx));
}
assert_eq!(ledger.remaining(), 3);
let (witnesses, _) = s.drain_all();
recover_drained_witnesses(witnesses, &h, &ledger);
assert_eq!(ledger.remaining(), 0, "a drained witness's unit must not stay outstanding");
h.reset().expect("pools whole");
}
#[test]
fn recovery_does_not_settle_a_unit_it_does_not_own() {
let h = handler();
let owner = DeviceCompletions::new().acquire(null_ptr());
let ledger = owner.ledger();
ledger.arm(7, ProofType::Basic.as_usize()).commit();
ledger.arm(7, ProofType::Compressor.as_usize()).commit();
assert_eq!(ledger.remaining(), 2);
let mut s = scheduler();
s.push(witness(&h, ProofType::Compressor, 7));
let (witnesses, _) = s.drain_all();
recover_drained_witnesses(witnesses, &h, &ledger);
assert_eq!(ledger.remaining(), 1, "only the compressor unit settles");
drop(ledger.adopt(7, ProofType::Basic.as_usize())); assert_eq!(ledger.remaining(), 0);
h.reset().expect("pools whole");
}
#[test]
fn a_witness_without_a_global_idx_still_returns_its_buffer() {
let h = handler();
let buf = h.take_buffer_witness();
let mut s = scheduler();
s.push(Proof::new_witness(ProofType::Recursive1, 0, 0, None, buf, 1));
let (witnesses, _) = s.drain_all();
recover_drained_witnesses(witnesses, &h, &DeviceCompletions::new().acquire(null_ptr()).ledger());
h.reset().expect("buffer returned even with no unit to settle");
}
#[test]
fn ledger_kind_spellings_agree() {
for (t, discriminant) in [
(ProofType::Basic, ProofType::Basic as usize),
(ProofType::Compressor, ProofType::Compressor as usize),
(ProofType::Recursive1, ProofType::Recursive1 as usize),
(ProofType::Recursive2, ProofType::Recursive2 as usize),
(ProofType::VadcopFinal, ProofType::VadcopFinal as usize),
(ProofType::VadcopFinalCompressed, ProofType::VadcopFinalCompressed as usize),
(ProofType::RecursiveF, ProofType::RecursiveF as usize),
(ProofType::RecurserAggregator, ProofType::RecurserAggregator as usize),
] {
assert_eq!(t.as_usize(), discriminant, "as_usize() must match `{t:?} as usize`");
}
}
#[test]
fn dropping_the_scheduler_undrained_is_what_loses_the_buffers() {
let h = handler();
let mut s = scheduler();
s.push(witness(&h, ProofType::Compressor, 0));
s.queues.clear(); assert!(h.reset().is_err(), "a dropped witness permanently shrinks its pool");
let _ = F::ZERO; }
}
#[cfg(test)]
mod admission_tests {
use super::{eligible_stream_count, witness_slot_cap};
const CLASSES: [usize; 3] = [7598, 6390, 6390];
#[test]
fn eligibility_matches_the_measured_carve() {
assert_eq!(eligible_stream_count(7598, &CLASSES), 1, "7.42 GB air: large class only");
assert_eq!(eligible_stream_count(6133, &CLASSES), 3, "5.99 GB air: all three");
assert_eq!(eligible_stream_count(9000, &CLASSES), 0, "bigger than any class");
}
#[test]
fn a_confined_air_gets_one_slot_per_stream_plus_a_prefetch() {
assert_eq!(witness_slot_cap(1, 3), 2);
assert_eq!(witness_slot_cap(2, 3), 3);
}
#[test]
fn an_air_that_fits_everywhere_is_uncapped() {
assert_eq!(witness_slot_cap(3, 3), usize::MAX);
assert_eq!(witness_slot_cap(8, 8), usize::MAX);
}
#[test]
fn unknown_eligibility_never_blocks() {
assert_eq!(eligible_stream_count(6133, &[]), 0);
assert_eq!(witness_slot_cap(0, 3), usize::MAX);
assert_eq!(witness_slot_cap(1, 0), usize::MAX);
}
}
#[cfg(test)]
mod schedule_tests {
use super::schedule_key;
use std::cmp::Reverse;
fn sort_ids(items: &[(usize, usize, bool, bool, u64)]) -> Vec<usize> {
let mut ids: Vec<usize> = (0..items.len()).collect();
ids.sort_by_key(|&i| {
let (ag, air, st, hc, cost) = items[i];
schedule_key(ag, air, st, hc, cost)
});
ids
}
#[test]
fn priority_tier_orders_stored_and_compressor_first() {
let items = [
(9, 0, false, false, 100), (8, 0, true, true, 100), (7, 0, false, true, 100), (6, 0, true, false, 100), ];
let tiers: Vec<u8> = sort_ids(&items)
.iter()
.map(|&i| schedule_key(items[i].0, items[i].1, items[i].2, items[i].3, items[i].4).0)
.collect();
assert_eq!(tiers, vec![0, 1, 2, 3]);
}
#[test]
fn lpt_orders_heavier_group_first_within_tier() {
let items = [(0, 0, false, false, 50), (0, 1, false, false, 500)];
let ordered = sort_ids(&items);
assert_eq!(ordered, vec![1, 0]); }
#[test]
fn key_is_reverse_on_cost() {
let hi = schedule_key(0, 0, false, false, 500);
let lo = schedule_key(0, 0, false, false, 50);
assert!(hi < lo); assert_eq!(hi.1, Reverse(500));
}
}