use std::collections::HashSet;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex, RwLock};
use std::time::Duration;
use crossbeam_channel::{unbounded, Receiver};
use proofman_starks_lib_c::{
clear_proof_done_callback_c, get_stream_proofs_c, get_stream_proofs_non_blocking_c, register_proof_done_callback_c,
CompletionMsg,
};
use crate::CancellationInfo;
const POLL_INTERVAL: Duration = Duration::from_micros(100);
const DRAIN_BUDGET: Duration = Duration::from_secs(5);
const SLOT_WAIT_WARN_AFTER: Duration = Duration::from_secs(10);
type UnitKey = (u64, usize);
pub struct DeviceCompletions {
slot: Arc<Mutex<bool>>,
next_epoch: AtomicU64,
}
impl Default for DeviceCompletions {
fn default() -> Self {
Self::new()
}
}
impl DeviceCompletions {
pub fn new() -> Self {
Self { slot: Arc::new(Mutex::new(false)), next_epoch: AtomicU64::new(1) }
}
pub fn acquire(&self, d_buffers: DeviceBuffersPtr) -> CompletionOwner {
let slot = SlotToken::take(&self.slot);
let epoch = self.next_epoch.fetch_add(1, Ordering::Relaxed);
let ledger = Arc::new(Ledger::new(epoch));
let (tx, rx) = unbounded::<CompletionMsg>();
register_proof_done_callback_c(tx);
CompletionOwner { ledger, rx, d_buffers, _slot: slot }
}
}
#[derive(Clone, Copy)]
pub struct DeviceBuffersPtr(pub *mut std::ffi::c_void);
unsafe impl Send for DeviceBuffersPtr {}
unsafe impl Sync for DeviceBuffersPtr {}
struct SlotToken {
slot: Arc<Mutex<bool>>,
}
impl SlotToken {
fn take(slot: &Arc<Mutex<bool>>) -> Self {
let start = std::time::Instant::now();
let mut warned = false;
loop {
{
let mut held = slot.lock().unwrap_or_else(|p| p.into_inner());
if !*held {
*held = true;
return Self { slot: Arc::clone(slot) };
}
}
if !warned && start.elapsed() >= SLOT_WAIT_WARN_AFTER {
warned = true;
tracing::warn!(
"Waiting >{}s for the proof-done completion capability: a previous CompletionOwner \
is still alive. Stop the outer-aggregation service before acquiring.",
SLOT_WAIT_WARN_AFTER.as_secs()
);
}
std::thread::sleep(POLL_INTERVAL);
}
}
}
impl Drop for SlotToken {
fn drop(&mut self) {
let mut held = self.slot.lock().unwrap_or_else(|p| p.into_inner());
*held = false;
}
}
pub struct Ledger {
epoch: u64,
outstanding: Mutex<HashSet<UnitKey>>,
remaining: AtomicUsize,
wait_lock: Mutex<()>,
cvar: Condvar,
}
impl Ledger {
fn new(epoch: u64) -> Self {
Self {
epoch,
outstanding: Mutex::new(HashSet::new()),
remaining: AtomicUsize::new(0),
wait_lock: Mutex::new(()),
cvar: Condvar::new(),
}
}
pub fn epoch(&self) -> u64 {
self.epoch
}
pub fn remaining(&self) -> usize {
self.remaining.load(Ordering::Acquire)
}
#[must_use = "dropping the token settles the unit immediately; hold it until the work is launched"]
pub fn arm(self: &Arc<Self>, id: u64, kind: usize) -> ProofToken {
let key = (id, kind);
{
let mut set = self.outstanding.lock().unwrap_or_else(|p| p.into_inner());
if set.insert(key) {
self.remaining.fetch_add(1, Ordering::AcqRel);
}
}
ProofToken { ledger: Arc::clone(self), key, armed: true }
}
#[must_use = "dropping the token settles the unit immediately; hold it until the work is launched"]
pub fn adopt(self: &Arc<Self>, id: u64, kind: usize) -> ProofToken {
ProofToken { ledger: Arc::clone(self), key: (id, kind), armed: true }
}
pub fn settle(&self, id: u64, kind: usize) {
let now = {
let mut set = self.outstanding.lock().unwrap_or_else(|p| p.into_inner());
if !set.remove(&(id, kind)) {
return;
}
self.remaining.fetch_sub(1, Ordering::AcqRel) - 1
};
if now == 0 {
let _g = self.wait_lock.lock().unwrap_or_else(|p| p.into_inner());
self.cvar.notify_all();
}
}
pub fn wait_settled<P: FnMut()>(
&self,
mut pump: P,
cancellation_info: &RwLock<CancellationInfo>,
timeout: Option<Duration>,
) -> bool {
let start = std::time::Instant::now();
let mut guard = self.wait_lock.lock().unwrap_or_else(|p| p.into_inner());
loop {
if self.remaining.load(Ordering::Acquire) == 0 {
return true;
}
let cancelled = {
let info = cancellation_info.read().unwrap_or_else(|p| p.into_inner());
info.token.is_cancelled()
};
if cancelled {
return false;
}
if let Some(limit) = timeout {
if start.elapsed() >= limit {
return false;
}
}
pump();
let (g, _) = self.cvar.wait_timeout(guard, POLL_INTERVAL).unwrap_or_else(|p| p.into_inner());
guard = g;
}
}
}
pub struct CompletionOwner {
ledger: Arc<Ledger>,
rx: Receiver<CompletionMsg>,
d_buffers: DeviceBuffersPtr,
_slot: SlotToken,
}
impl CompletionOwner {
pub fn ledger(&self) -> Arc<Ledger> {
Arc::clone(&self.ledger)
}
pub fn receiver(&self) -> Receiver<CompletionMsg> {
self.rx.clone()
}
pub fn epoch(&self) -> u64 {
self.ledger.epoch()
}
pub fn remaining(&self) -> usize {
self.ledger.remaining()
}
pub fn wait_settled<P: FnMut()>(
&self,
pump: P,
cancellation_info: &RwLock<CancellationInfo>,
timeout: Option<Duration>,
) -> bool {
self.ledger.wait_settled(pump, cancellation_info, timeout)
}
}
impl Drop for CompletionOwner {
fn drop(&mut self) {
let start = std::time::Instant::now();
while self.ledger.remaining() > 0 && start.elapsed() < DRAIN_BUDGET {
if !self.d_buffers.0.is_null() {
get_stream_proofs_non_blocking_c(self.d_buffers.0);
}
std::thread::sleep(POLL_INTERVAL);
}
if !self.d_buffers.0.is_null() {
get_stream_proofs_c(self.d_buffers.0);
}
let leaked = self.ledger.remaining();
if leaked > 0 {
tracing::debug!(
"CompletionOwner(epoch {}) released with {} unit(s) unsettled (expected after cancellation)",
self.ledger.epoch(),
leaked
);
}
clear_proof_done_callback_c();
}
}
#[must_use = "dropping the token settles the unit immediately; hold it until the work is launched"]
pub struct ProofToken {
ledger: Arc<Ledger>,
key: UnitKey,
armed: bool,
}
impl ProofToken {
pub fn commit(mut self) {
self.armed = false;
}
pub fn id(&self) -> u64 {
self.key.0
}
}
impl Drop for ProofToken {
fn drop(&mut self) {
if self.armed {
self.ledger.settle(self.key.0, self.key.1);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const BASIC: usize = 0;
const RECURSIVE1: usize = 2;
fn owner() -> CompletionOwner {
DeviceCompletions::new().acquire(DeviceBuffersPtr(std::ptr::null_mut()))
}
#[test]
fn arm_then_drop_settles() {
let o = owner();
let l = o.ledger();
let t = l.arm(7, BASIC);
assert_eq!(l.remaining(), 1);
drop(t);
assert_eq!(l.remaining(), 0);
}
#[test]
fn commit_transfers_settlement_to_the_async_path() {
let o = owner();
let l = o.ledger();
l.arm(7, BASIC).commit();
assert_eq!(l.remaining(), 1, "a committed unit stays outstanding until it completes");
l.settle(7, BASIC);
assert_eq!(l.remaining(), 0);
}
#[test]
fn adopt_does_not_change_the_count_but_settles_on_drop() {
let o = owner();
let l = o.ledger();
l.arm(7, BASIC).commit();
let taken = l.adopt(7, BASIC);
assert_eq!(l.remaining(), 1, "adopt must not double-count an already-armed unit");
drop(taken);
assert_eq!(l.remaining(), 0, "an uncommitted adopt settles the unit it took over");
}
#[test]
fn settling_an_unknown_unit_is_a_no_op() {
let o = owner();
let l = o.ledger();
l.settle(1234, BASIC);
assert_eq!(l.remaining(), 0);
l.arm(1, BASIC).commit();
l.settle(999, BASIC);
assert_eq!(l.remaining(), 1, "a stray completion must not settle someone else's unit");
}
#[test]
fn the_same_id_with_a_different_kind_is_a_distinct_unit() {
let o = owner();
let l = o.ledger();
l.arm(5, BASIC).commit();
l.arm(5, RECURSIVE1).commit();
assert_eq!(l.remaining(), 2, "(5, Basic) and (5, Recursive1) are different units");
l.settle(5, BASIC);
assert_eq!(l.remaining(), 1, "settling the basic unit must not settle the recursive one");
l.settle(5, RECURSIVE1);
assert_eq!(l.remaining(), 0);
}
#[test]
fn settling_twice_only_counts_once() {
let o = owner();
let l = o.ledger();
l.arm(3, BASIC).commit();
l.settle(3, BASIC);
l.settle(3, BASIC);
assert_eq!(l.remaining(), 0);
}
#[test]
fn a_panicking_worker_still_settles_its_unit() {
let o = owner();
let l = o.ledger();
let l2 = Arc::clone(&l);
let r = std::panic::catch_unwind(std::panic::AssertUnwindSafe(move || {
let _t = l2.arm(11, BASIC);
panic!("failure between arming and launch");
}));
assert!(r.is_err());
assert_eq!(l.remaining(), 0);
}
#[test]
fn workers_learn_to_exit_when_the_owner_drops() {
let o = owner();
let rx = o.receiver();
drop(o);
assert!(rx.recv().is_err(), "dropping the owner disconnects consumers; no sentinel needed");
}
#[test]
fn the_capability_is_exclusive() {
let completions = DeviceCompletions::new();
let null = DeviceBuffersPtr(std::ptr::null_mut());
let first = completions.acquire(null);
let epoch = first.epoch();
drop(first);
let second = completions.acquire(null);
assert!(second.epoch() > epoch, "each owner gets a fresh epoch");
}
}