use std::cell::RefCell;
use std::fmt;
use std::sync::Arc;
use std::time::{Duration, Instant};
use cudarc::driver::sys::CUdeviceptr;
use onnx_runtime_ep_api::{LazyWeight, MmapRegionSource};
use crate::pinned_pool::PinnedStagingPool;
use crate::runtime::{CudaRuntime, FailedHtodCompletion, PinnedStaging};
use crate::weight_paging::fill_staging_from_regions;
pub type FenceId = u64;
const SLOTS: usize = 2;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum PrefillSlotStatus {
Free,
Filling,
Ready,
InUse,
Draining,
Poisoned,
}
impl PrefillSlotStatus {
fn is_claimable(self) -> bool {
matches!(self, Self::Free | Self::Draining)
}
}
#[derive(Debug)]
pub enum PrefillReject<E> {
PoolCapacity { layer_bytes: u64 },
CaptureActive,
EmptyLayer,
SlotsBusy,
StaleGeneration { layer_id: u64 },
WrongState {
layer_id: u64,
status: PrefillSlotStatus,
},
Poisoned { layer_id: u64 },
Disabled,
Transfer(E),
}
impl<E: fmt::Display> fmt::Display for PrefillReject<E> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::PoolCapacity { layer_bytes } => write!(
f,
"prefill double buffer declined: pinned pool cannot retain {SLOTS} concurrent \
buffers of {layer_bytes} bytes; fall back to synchronous single-buffer prefill"
),
Self::CaptureActive => write!(
f,
"prefill double buffer declined: CUDA graph capture/replay active, buffer \
reservation must run at a scheduler safe point"
),
Self::EmptyLayer => {
write!(
f,
"prefill double buffer declined: layer has zero transferable bytes"
)
}
Self::SlotsBusy => write!(
f,
"prefill double buffer declined: both slots occupied by not-yet-released layers \
(pipeline depth limit; consume and release before prefetching further)"
),
Self::StaleGeneration { layer_id } => write!(
f,
"prefill double buffer refused a stale operation for layer {layer_id}: its slot \
was already reused for another layer"
),
Self::WrongState { layer_id, status } => write!(
f,
"prefill double buffer refused an operation for layer {layer_id}: slot is {status:?}"
),
Self::Poisoned { layer_id } => write!(
f,
"prefill double buffer refused an operation for layer {layer_id}: slot is poisoned"
),
Self::Disabled => write!(
f,
"prefill double buffer is disabled (default-off): set \
ONNX_GENAI_PREFILL_DOUBLE_BUFFER=1 to enable, or use the synchronous \
single-buffer prefill path"
),
Self::Transfer(error) => write!(f, "prefill double buffer transfer failed: {error}"),
}
}
}
#[derive(Clone, Copy, Debug, Default)]
pub struct SlotFillPlan {
pub prev_copy_fence: Option<FenceId>,
pub prev_release_fence: Option<FenceId>,
}
#[derive(Clone, Copy, Debug)]
pub struct FillOutcome {
pub copy_fence: FenceId,
pub reuse_wait_ns: u64,
}
#[derive(Clone, Copy, Debug, Default)]
pub struct SlotTeardown {
pub last_copy_fence: Option<FenceId>,
pub last_release_fence: Option<FenceId>,
pub live: bool,
}
pub trait PrefillTransfer {
type Payload: Clone;
type LayerReq<'a>;
type Error: fmt::Display;
fn layer_bytes(&self, req: &Self::LayerReq<'_>) -> u64;
fn capture_active(&self) -> bool;
fn can_retain_concurrent(&self, layer_bytes: u64) -> bool;
fn reserve(&mut self, layer_bytes: u64) -> Result<(), Self::Error>;
fn fill_slot(
&self,
slot: usize,
req: &Self::LayerReq<'_>,
plan: SlotFillPlan,
) -> Result<FillOutcome, Self::Error>;
fn payload(&self, slot: usize) -> Self::Payload;
fn compute_wait(&self, copy_fence: FenceId) -> Result<(), Self::Error>;
fn record_release_fence(&self) -> Result<FenceId, Self::Error>;
fn quarantine_slot(&self, slot: usize);
fn teardown(&mut self, slots: [SlotTeardown; SLOTS]);
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct LayerTicket {
layer_id: u64,
slot: usize,
generation: u64,
}
impl LayerTicket {
pub fn layer_id(&self) -> u64 {
self.layer_id
}
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct PrefillMetrics {
pub layers_prefetched: u64,
pub layers_consumed: u64,
pub layers_released: u64,
pub reuse_wait_ns: u64,
pub declined_slots_busy: u64,
pub declined_capture: u64,
pub declined_empty: u64,
pub stale_rejected: u64,
pub cancelled: u64,
pub poisoned: u64,
}
#[derive(Debug)]
struct Slot {
status: PrefillSlotStatus,
generation: u64,
layer_id: u64,
copy_fence: Option<FenceId>,
release_fence: Option<FenceId>,
}
impl Default for Slot {
fn default() -> Self {
Self {
status: PrefillSlotStatus::Free,
generation: 0,
layer_id: 0,
copy_fence: None,
release_fence: None,
}
}
}
pub struct PrefillDoubleBuffer<T: PrefillTransfer> {
transfer: T,
slots: [Slot; SLOTS],
layer_bytes: u64,
metrics: PrefillMetrics,
torn_down: bool,
}
impl<T: PrefillTransfer> fmt::Debug for PrefillDoubleBuffer<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PrefillDoubleBuffer")
.field("slots", &self.slots)
.field("layer_bytes", &self.layer_bytes)
.field("metrics", &self.metrics)
.field("torn_down", &self.torn_down)
.finish_non_exhaustive()
}
}
impl<T: PrefillTransfer> PrefillDoubleBuffer<T> {
pub fn new(mut transfer: T, layer_bytes: u64) -> Result<Self, PrefillReject<T::Error>> {
if layer_bytes == 0 {
return Err(PrefillReject::EmptyLayer);
}
if transfer.capture_active() {
return Err(PrefillReject::CaptureActive);
}
if !transfer.can_retain_concurrent(layer_bytes) {
return Err(PrefillReject::PoolCapacity { layer_bytes });
}
transfer
.reserve(layer_bytes)
.map_err(PrefillReject::Transfer)?;
Ok(Self {
transfer,
slots: Default::default(),
layer_bytes,
metrics: PrefillMetrics::default(),
torn_down: false,
})
}
pub fn layer_bytes(&self) -> u64 {
self.layer_bytes
}
pub fn metrics(&self) -> PrefillMetrics {
self.metrics
}
pub fn slot_status(&self, idx: usize) -> PrefillSlotStatus {
self.slots[idx].status
}
pub fn transfer(&self) -> &T {
&self.transfer
}
pub fn prefetch(
&mut self,
layer_id: u64,
req: &T::LayerReq<'_>,
) -> Result<LayerTicket, PrefillReject<T::Error>> {
debug_assert!(!self.torn_down, "prefetch after teardown");
if self.transfer.capture_active() {
self.metrics.declined_capture += 1;
return Err(PrefillReject::CaptureActive);
}
if self.transfer.layer_bytes(req) == 0 {
self.metrics.declined_empty += 1;
return Err(PrefillReject::EmptyLayer);
}
let Some(slot_idx) = self.pick_claimable_slot() else {
self.metrics.declined_slots_busy += 1;
return Err(PrefillReject::SlotsBusy);
};
let reused = self.slots[slot_idx].status == PrefillSlotStatus::Draining;
let plan = if reused {
debug_assert!(
self.slots[slot_idx].release_fence.is_some(),
"reuse-before-release: a Draining slot must carry a release fence"
);
SlotFillPlan {
prev_copy_fence: self.slots[slot_idx].copy_fence,
prev_release_fence: self.slots[slot_idx].release_fence,
}
} else {
SlotFillPlan::default()
};
let generation = self.slots[slot_idx].generation + 1;
self.slots[slot_idx].status = PrefillSlotStatus::Filling;
self.slots[slot_idx].generation = generation;
self.slots[slot_idx].layer_id = layer_id;
match self.transfer.fill_slot(slot_idx, req, plan) {
Ok(outcome) => {
self.slots[slot_idx].copy_fence = Some(outcome.copy_fence);
self.slots[slot_idx].release_fence = None;
self.slots[slot_idx].status = PrefillSlotStatus::Ready;
self.metrics.layers_prefetched += 1;
self.metrics.reuse_wait_ns = self
.metrics
.reuse_wait_ns
.saturating_add(outcome.reuse_wait_ns);
Ok(LayerTicket {
layer_id,
slot: slot_idx,
generation,
})
}
Err(error) => {
self.poison(slot_idx);
Err(PrefillReject::Transfer(error))
}
}
}
pub fn wait(&mut self, ticket: &LayerTicket) -> Result<T::Payload, PrefillReject<T::Error>> {
self.validate_ticket(ticket)?;
let slot_idx = ticket.slot;
match self.slots[slot_idx].status {
PrefillSlotStatus::Ready => {}
PrefillSlotStatus::Poisoned => {
return Err(PrefillReject::Poisoned {
layer_id: ticket.layer_id,
});
}
status => {
return Err(PrefillReject::WrongState {
layer_id: ticket.layer_id,
status,
});
}
}
let copy_fence = self.slots[slot_idx].copy_fence.unwrap_or(0);
if let Err(error) = self.transfer.compute_wait(copy_fence) {
self.poison(slot_idx);
return Err(PrefillReject::Transfer(error));
}
self.slots[slot_idx].status = PrefillSlotStatus::InUse;
self.metrics.layers_consumed += 1;
Ok(self.transfer.payload(slot_idx))
}
pub fn release(&mut self, ticket: LayerTicket) -> Result<(), PrefillReject<T::Error>> {
self.validate_ticket(&ticket)?;
let slot_idx = ticket.slot;
match self.slots[slot_idx].status {
PrefillSlotStatus::InUse => {}
PrefillSlotStatus::Poisoned => {
return Err(PrefillReject::Poisoned {
layer_id: ticket.layer_id,
});
}
status => {
return Err(PrefillReject::WrongState {
layer_id: ticket.layer_id,
status,
});
}
}
self.record_release(slot_idx)?;
self.metrics.layers_released += 1;
Ok(())
}
pub fn cancel(&mut self, ticket: LayerTicket) -> Result<(), PrefillReject<T::Error>> {
self.validate_ticket(&ticket)?;
let slot_idx = ticket.slot;
match self.slots[slot_idx].status {
PrefillSlotStatus::Ready | PrefillSlotStatus::InUse => {}
PrefillSlotStatus::Poisoned => {
return Err(PrefillReject::Poisoned {
layer_id: ticket.layer_id,
});
}
status => {
return Err(PrefillReject::WrongState {
layer_id: ticket.layer_id,
status,
});
}
}
self.record_release(slot_idx)?;
self.metrics.cancelled += 1;
Ok(())
}
fn pick_claimable_slot(&self) -> Option<usize> {
(0..SLOTS)
.find(|&i| self.slots[i].status == PrefillSlotStatus::Free)
.or_else(|| (0..SLOTS).find(|&i| self.slots[i].status.is_claimable()))
}
fn record_release(&mut self, slot_idx: usize) -> Result<(), PrefillReject<T::Error>> {
match self.transfer.record_release_fence() {
Ok(fence) => {
self.slots[slot_idx].release_fence = Some(fence);
self.slots[slot_idx].status = PrefillSlotStatus::Draining;
Ok(())
}
Err(error) => {
self.poison(slot_idx);
Err(PrefillReject::Transfer(error))
}
}
}
fn validate_ticket(&mut self, ticket: &LayerTicket) -> Result<(), PrefillReject<T::Error>> {
if self.slots[ticket.slot].generation != ticket.generation {
self.metrics.stale_rejected += 1;
return Err(PrefillReject::StaleGeneration {
layer_id: ticket.layer_id,
});
}
Ok(())
}
fn poison(&mut self, slot_idx: usize) {
self.transfer.quarantine_slot(slot_idx);
self.slots[slot_idx].status = PrefillSlotStatus::Poisoned;
self.metrics.poisoned += 1;
}
}
impl<T: PrefillTransfer> Drop for PrefillDoubleBuffer<T> {
fn drop(&mut self) {
if self.torn_down {
return;
}
self.torn_down = true;
let descriptors = std::array::from_fn(|i| {
let slot = &self.slots[i];
let live = !matches!(slot.status, PrefillSlotStatus::Poisoned);
SlotTeardown {
last_copy_fence: if live { slot.copy_fence } else { None },
last_release_fence: if live { slot.release_fence } else { None },
live,
}
});
self.transfer.teardown(descriptors);
}
}
pub fn duration_ns(d: Duration) -> u64 {
u64::try_from(d.as_nanos()).unwrap_or(u64::MAX)
}
#[derive(Debug)]
pub struct CudaPrefillError(String);
impl CudaPrefillError {
fn new(msg: impl Into<String>) -> Self {
Self(msg.into())
}
}
impl fmt::Display for CudaPrefillError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl std::error::Error for CudaPrefillError {}
pub struct PrefillLayerRequest<'a> {
pub weight: &'a LazyWeight,
pub source: &'a dyn MmapRegionSource,
}
impl<'a> PrefillLayerRequest<'a> {
pub fn new(weight: &'a LazyWeight, source: &'a dyn MmapRegionSource) -> Self {
Self { weight, source }
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct PrefillLayerView {
pub device_ptr: CUdeviceptr,
pub len: usize,
}
struct CudaSlotBuffers {
device_ptr: CUdeviceptr,
capacity: usize,
valid_len: usize,
staging: Option<PinnedStaging>,
reuse_fence: FenceId,
}
impl CudaSlotBuffers {
fn empty() -> Self {
Self {
device_ptr: 0,
capacity: 0,
valid_len: 0,
staging: None,
reuse_fence: 0,
}
}
}
pub struct CudaPrefillTransfer {
runtime: Arc<CudaRuntime>,
staging_pool: Arc<PinnedStagingPool>,
slots: [RefCell<CudaSlotBuffers>; SLOTS],
quarantined: RefCell<Vec<CudaSlotBuffers>>,
}
impl CudaPrefillTransfer {
pub fn new(runtime: Arc<CudaRuntime>, staging_pool: Arc<PinnedStagingPool>) -> Self {
Self {
runtime,
staging_pool,
slots: std::array::from_fn(|_| RefCell::new(CudaSlotBuffers::empty())),
quarantined: RefCell::new(Vec::new()),
}
}
pub fn quarantined_len(&self) -> usize {
self.quarantined.borrow().len()
}
}
impl PrefillTransfer for CudaPrefillTransfer {
type Payload = PrefillLayerView;
type LayerReq<'a> = PrefillLayerRequest<'a>;
type Error = CudaPrefillError;
fn layer_bytes(&self, req: &Self::LayerReq<'_>) -> u64 {
req.weight.region_bytes_len() as u64
}
fn capture_active(&self) -> bool {
self.runtime.is_capturing().unwrap_or(true)
}
fn can_retain_concurrent(&self, layer_bytes: u64) -> bool {
self.staging_pool
.can_retain_concurrent(layer_bytes as usize, SLOTS)
}
fn reserve(&mut self, layer_bytes: u64) -> Result<(), Self::Error> {
let len = layer_bytes as usize;
let mut staging_bufs: Vec<PinnedStaging> = Vec::with_capacity(SLOTS);
let mut device_ptrs: Vec<CUdeviceptr> = Vec::with_capacity(SLOTS);
let mut acquire = || -> Result<(), Self::Error> {
for _ in 0..SLOTS {
let pooled = self.staging_pool.acquire(len).map_err(|error| {
CudaPrefillError::new(format!("pinned staging acquire: {error}"))
})?;
staging_bufs.push(pooled.into_inner());
let ptr = self
.runtime
.alloc_raw(len)
.map_err(|error| CudaPrefillError::new(format!("device alloc: {error}")))?;
device_ptrs.push(ptr);
}
Ok(())
};
if let Err(error) = acquire() {
for ptr in device_ptrs {
let _ = unsafe { self.runtime.free_raw(ptr) };
}
return Err(error);
}
for (i, (ptr, staging)) in device_ptrs.into_iter().zip(staging_bufs).enumerate() {
let mut slot = self.slots[i].borrow_mut();
slot.device_ptr = ptr;
slot.capacity = len;
slot.valid_len = 0;
slot.staging = Some(staging);
slot.reuse_fence = 0;
}
Ok(())
}
fn fill_slot(
&self,
slot: usize,
req: &Self::LayerReq<'_>,
plan: SlotFillPlan,
) -> Result<FillOutcome, Self::Error> {
let mut s = self.slots[slot].borrow_mut();
let mut reuse_wait_ns = 0u64;
if plan.prev_copy_fence.is_some() {
let fence = std::mem::take(&mut s.reuse_fence);
let start = Instant::now();
if let Err(error) = self.runtime.resolve_prefetch_fence(fence) {
let (detail, completion) = error.into_parts();
if matches!(completion, FailedHtodCompletion::MayBeInFlight) {
return Err(CudaPrefillError::new(format!(
"reuse drain could not establish prior copy completion: {detail}"
)));
}
}
reuse_wait_ns = duration_ns(start.elapsed());
}
let src_len = req.weight.region_bytes_len();
if src_len == 0 {
return Err(CudaPrefillError::new("fill_slot on a zero-byte layer"));
}
if src_len > s.capacity {
return Err(CudaPrefillError::new(format!(
"layer is {src_len} bytes but the pipeline reserved {}-byte buffers",
s.capacity
)));
}
{
let staging = s
.staging
.as_mut()
.ok_or_else(|| CudaPrefillError::new("fill_slot on an unreserved slot"))?;
fill_staging_from_regions(req.weight, req.source, staging)
.map_err(|error| CudaPrefillError::new(format!("staging fill: {error}")))?;
}
if let Some(rel) = plan.prev_release_fence {
self.runtime
.copy_wait_fence(rel)
.map_err(|error| CudaPrefillError::new(format!("copy_wait release: {error}")))?;
}
let ptr = s.device_ptr;
let enqueue = {
let staging = s.staging.as_ref().expect("staging present after fill");
unsafe { self.runtime.htod_async(&staging.as_slice()[..src_len], ptr) }
};
if let Err(error) = enqueue {
return Err(CudaPrefillError::new(format!("H2D enqueue: {error}")));
}
s.valid_len = src_len;
let copy_fence = self
.runtime
.record_copy_fence()
.map_err(|error| CudaPrefillError::new(format!("record copy fence: {error}")))?;
let reuse_fence = self
.runtime
.record_copy_fence()
.map_err(|error| CudaPrefillError::new(format!("record reuse fence: {error}")))?;
s.reuse_fence = reuse_fence;
Ok(FillOutcome {
copy_fence,
reuse_wait_ns,
})
}
fn payload(&self, slot: usize) -> Self::Payload {
let s = self.slots[slot].borrow();
PrefillLayerView {
device_ptr: s.device_ptr,
len: s.valid_len,
}
}
fn compute_wait(&self, copy_fence: FenceId) -> Result<(), Self::Error> {
self.runtime
.compute_wait_fence(copy_fence)
.map_err(|error| CudaPrefillError::new(format!("compute_wait: {error}")))
}
fn record_release_fence(&self) -> Result<FenceId, Self::Error> {
self.runtime
.record_compute_fence()
.map_err(|error| CudaPrefillError::new(format!("record release fence: {error}")))
}
fn quarantine_slot(&self, slot: usize) {
let taken = std::mem::replace(
&mut *self.slots[slot].borrow_mut(),
CudaSlotBuffers::empty(),
);
self.quarantined.borrow_mut().push(taken);
}
fn teardown(&mut self, slots: [SlotTeardown; SLOTS]) {
let _ = self.runtime.drain_for_unmap();
for (i, descriptor) in slots.iter().enumerate() {
if !descriptor.live {
continue;
}
let mut s = self.slots[i].borrow_mut();
let reuse_fence = std::mem::take(&mut s.reuse_fence);
let witness = match self.runtime.resolve_prefetch_fence(reuse_fence) {
Ok(completed) => Some(completed),
Err(error) => {
let (_detail, completion) = error.into_parts();
if matches!(completion, FailedHtodCompletion::MayBeInFlight) {
let taken = std::mem::replace(&mut *s, CudaSlotBuffers::empty());
drop(s);
self.quarantined.borrow_mut().push(taken);
continue;
}
None
}
};
let _ = self
.runtime
.resolve_prefetch_fence(descriptor.last_copy_fence.unwrap_or(0));
let _ = self
.runtime
.resolve_prefetch_fence(descriptor.last_release_fence.unwrap_or(0));
let ptr = std::mem::replace(&mut s.device_ptr, 0);
if ptr != 0 {
let _ = unsafe { self.runtime.free_raw(ptr) };
}
if let (Some(staging), Some(witness)) = (s.staging.take(), witness) {
self.staging_pool.release(staging, witness);
}
}
}
}
impl Drop for CudaPrefillTransfer {
fn drop(&mut self) {
for taken in self.quarantined.get_mut().drain(..) {
std::mem::forget(taken);
}
for slot in &self.slots {
let mut s = slot.borrow_mut();
let ptr = std::mem::replace(&mut s.device_ptr, 0);
if ptr != 0 {
let _ = unsafe { self.runtime.free_raw(ptr) };
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::cell::RefCell;
use std::rc::Rc;
#[derive(Clone, Debug, PartialEq, Eq)]
enum Op {
Reserve {
bytes: u64,
},
ReuseWait {
slot: usize,
fence: FenceId,
},
CopyWait {
slot: usize,
fence: FenceId,
},
Fill {
slot: usize,
layer: u64,
copy_fence: FenceId,
},
ComputeWait {
fence: FenceId,
},
ReleaseFence {
fence: FenceId,
},
Quarantine {
slot: usize,
},
Teardown {
live: [bool; SLOTS],
fences: [Option<FenceId>; SLOTS],
},
}
#[derive(Clone, Debug)]
struct FakeLayer {
id: u64,
bytes: u64,
}
#[derive(Default)]
struct FakeState {
ops: Vec<Op>,
next_fence: FenceId,
slot_layer: [Option<u64>; SLOTS],
buffer_live: [bool; SLOTS],
reserved: bool,
capture: bool,
pool_ok: bool,
reserve_fails: bool,
fail_fill_layers: Vec<u64>,
fail_compute_wait: bool,
fail_release_fence: bool,
reuse_wait_ns: u64,
}
impl FakeState {
fn fence(&mut self) -> FenceId {
self.next_fence += 1;
self.next_fence
}
}
#[derive(Clone)]
struct FakeTransfer {
state: Rc<RefCell<FakeState>>,
}
#[derive(Debug)]
struct FakeError(String);
impl fmt::Display for FakeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.0)
}
}
impl FakeTransfer {
fn new() -> Self {
Self {
state: Rc::new(RefCell::new(FakeState {
pool_ok: true,
..Default::default()
})),
}
}
fn ops(&self) -> Vec<Op> {
self.state.borrow().ops.clone()
}
}
impl PrefillTransfer for FakeTransfer {
type Payload = u64;
type LayerReq<'a> = FakeLayer;
type Error = FakeError;
fn layer_bytes(&self, req: &Self::LayerReq<'_>) -> u64 {
req.bytes
}
fn capture_active(&self) -> bool {
self.state.borrow().capture
}
fn can_retain_concurrent(&self, _layer_bytes: u64) -> bool {
self.state.borrow().pool_ok
}
fn reserve(&mut self, layer_bytes: u64) -> Result<(), Self::Error> {
let mut s = self.state.borrow_mut();
if s.reserve_fails {
return Err(FakeError("reserve failed".into()));
}
s.ops.push(Op::Reserve { bytes: layer_bytes });
s.reserved = true;
s.buffer_live = [true, true];
Ok(())
}
fn fill_slot(
&self,
slot: usize,
req: &Self::LayerReq<'_>,
plan: SlotFillPlan,
) -> Result<FillOutcome, Self::Error> {
let mut s = self.state.borrow_mut();
let mut reuse_wait_ns = 0;
if let Some(prev) = plan.prev_copy_fence {
reuse_wait_ns = s.reuse_wait_ns;
s.ops.push(Op::ReuseWait { slot, fence: prev });
}
if let Some(rel) = plan.prev_release_fence {
s.ops.push(Op::CopyWait { slot, fence: rel });
}
if s.fail_fill_layers.contains(&req.id) {
return Err(FakeError(format!("fill failed for layer {}", req.id)));
}
let copy_fence = s.fence();
s.ops.push(Op::Fill {
slot,
layer: req.id,
copy_fence,
});
s.slot_layer[slot] = Some(req.id);
Ok(FillOutcome {
copy_fence,
reuse_wait_ns,
})
}
fn payload(&self, slot: usize) -> Self::Payload {
self.state.borrow().slot_layer[slot].expect("payload of a filled slot")
}
fn compute_wait(&self, copy_fence: FenceId) -> Result<(), Self::Error> {
let mut s = self.state.borrow_mut();
if s.fail_compute_wait {
return Err(FakeError("compute_wait failed".into()));
}
s.ops.push(Op::ComputeWait { fence: copy_fence });
Ok(())
}
fn record_release_fence(&self) -> Result<FenceId, Self::Error> {
let mut s = self.state.borrow_mut();
if s.fail_release_fence {
return Err(FakeError("release fence failed".into()));
}
let fence = s.fence();
s.ops.push(Op::ReleaseFence { fence });
Ok(fence)
}
fn quarantine_slot(&self, slot: usize) {
let mut s = self.state.borrow_mut();
s.ops.push(Op::Quarantine { slot });
s.buffer_live[slot] = false;
}
fn teardown(&mut self, slots: [SlotTeardown; SLOTS]) {
let mut s = self.state.borrow_mut();
let live = [slots[0].live, slots[1].live];
let fences = [slots[0].last_copy_fence, slots[1].last_copy_fence];
s.ops.push(Op::Teardown { live, fences });
for (i, d) in slots.iter().enumerate() {
if d.live {
let _ = (d.last_copy_fence, d.last_release_fence);
s.buffer_live[i] = false;
}
}
}
}
fn layer(id: u64) -> FakeLayer {
FakeLayer { id, bytes: 1 << 20 }
}
fn cycle(db: &mut PrefillDoubleBuffer<FakeTransfer>, id: u64) {
let ticket = db.prefetch(id, &layer(id)).expect("prefetch");
let payload = db.wait(&ticket).expect("wait");
assert_eq!(payload, id, "consumer read the layer it prefetched");
db.release(ticket).expect("release");
}
fn assert_no_leak(fake: &FakeTransfer) {
let s = fake.state.borrow();
assert_eq!(
s.buffer_live,
[false, false],
"every reserved buffer must be released or quarantined exactly once"
);
}
#[test]
fn reserves_two_buffers_then_runs_n_and_n_plus_one_in_order() {
let fake = FakeTransfer::new();
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
let t0 = db.prefetch(0, &layer(0)).unwrap();
let t1 = db.prefetch(1, &layer(1)).unwrap();
assert_ne!(t0.slot, t1.slot, "N and N+1 occupy different slots");
assert_eq!(db.wait(&t0).unwrap(), 0);
db.release(t0).unwrap();
assert_eq!(db.wait(&t1).unwrap(), 1);
db.release(t1).unwrap();
let ops = fake.ops();
let fill0 = ops.iter().find_map(|op| match op {
Op::Fill {
layer: 0,
copy_fence,
..
} => Some(*copy_fence),
_ => None,
});
assert!(
ops.contains(&Op::ComputeWait {
fence: fill0.unwrap()
}),
"wait(0) must order compute after fill(0)'s copy fence"
);
drop(db);
assert_no_leak(&fake);
}
#[test]
fn wraparound_reuses_both_slots_with_two_directional_fencing() {
let fake = FakeTransfer::new();
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
for id in 0..6 {
cycle(&mut db, id);
}
let ops = fake.ops();
let reuse_waits = ops
.iter()
.filter(|op| matches!(op, Op::ReuseWait { .. }))
.count();
let copy_waits = ops
.iter()
.filter(|op| matches!(op, Op::CopyWait { .. }))
.count();
assert_eq!(
reuse_waits, 4,
"layers 2..6 each reuse a slot (staging WAR host-wait)"
);
assert_eq!(
copy_waits, 4,
"layers 2..6 each reuse a slot (device WAR copy_wait)"
);
let idx_fill2 = ops
.iter()
.position(|op| matches!(op, Op::Fill { layer: 2, .. }))
.unwrap();
assert!(matches!(ops[idx_fill2 - 1], Op::CopyWait { slot: 0, .. }));
assert!(matches!(ops[idx_fill2 - 2], Op::ReuseWait { slot: 0, .. }));
assert_eq!(db.metrics().layers_prefetched, 6);
assert_eq!(db.metrics().layers_released, 6);
drop(db);
assert_no_leak(&fake);
}
#[test]
fn single_layer_prefill_never_reuses() {
let fake = FakeTransfer::new();
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
cycle(&mut db, 0);
let ops = fake.ops();
assert!(
!ops.iter()
.any(|op| matches!(op, Op::ReuseWait { .. } | Op::CopyWait { .. })),
"a single layer performs no reuse fencing"
);
assert_eq!(db.metrics().reuse_wait_ns, 0);
drop(db);
assert_no_leak(&fake);
}
#[test]
fn final_layer_release_then_teardown_drains_in_flight() {
let fake = FakeTransfer::new();
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
let t0 = db.prefetch(0, &layer(0)).unwrap();
let t1 = db.prefetch(1, &layer(1)).unwrap();
let _ = db.wait(&t0).unwrap();
db.release(t0).unwrap();
assert_eq!(db.slot_status(t1.slot), PrefillSlotStatus::Ready);
drop(db);
let ops = fake.ops();
let teardown = ops.iter().find_map(|op| match op {
Op::Teardown { live, fences } => Some((*live, *fences)),
_ => None,
});
let (live, fences) = teardown.expect("teardown ran");
assert!(live[t1.slot], "the in-flight slot is live at teardown");
assert!(
fences[t1.slot].is_some(),
"teardown drains the in-flight copy fence"
);
assert_no_leak(&fake);
}
#[test]
fn cancellation_midflight_frees_slot_for_reuse() {
let fake = FakeTransfer::new();
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
let t0 = db.prefetch(0, &layer(0)).unwrap();
db.cancel(t0.clone()).unwrap();
assert_eq!(db.metrics().cancelled, 1);
assert_eq!(db.slot_status(t0.slot), PrefillSlotStatus::Draining);
let t_other = db.prefetch(1, &layer(1)).unwrap();
assert_ne!(t_other.slot, t0.slot, "the fresh Free slot is taken first");
let t_new = db.prefetch(2, &layer(2)).unwrap();
assert_eq!(t_new.slot, t0.slot, "cancelled slot is reused");
assert_eq!(db.wait(&t_new).unwrap(), 2);
db.release(t_new).unwrap();
let _ = db.wait(&t_other).unwrap();
db.release(t_other).unwrap();
let ops = fake.ops();
assert!(
ops.iter()
.any(|op| matches!(op, Op::ReuseWait { slot, .. } if *slot == t0.slot)),
"reusing a cancelled slot drains its in-flight copy before refill"
);
drop(db);
assert_no_leak(&fake);
}
#[test]
fn stale_ticket_after_reuse_is_refused_not_served() {
let fake = FakeTransfer::new();
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
let stale = db.prefetch(0, &layer(0)).unwrap();
let _ = db.wait(&stale).unwrap();
db.release(stale.clone()).unwrap();
let other = db.prefetch(1, &layer(1)).unwrap();
let reuse = db.prefetch(2, &layer(2)).unwrap();
assert_eq!(reuse.slot, stale.slot, "layer 2 reused layer 0's slot");
match db.wait(&stale) {
Err(PrefillReject::StaleGeneration { layer_id: 0 }) => {}
other => panic!("expected StaleGeneration, got {other:?}"),
}
assert_eq!(db.metrics().stale_rejected, 1);
let _ = db.wait(&other).unwrap();
db.release(other).unwrap();
let _ = db.wait(&reuse).unwrap();
db.release(reuse).unwrap();
drop(db);
assert_no_leak(&fake);
}
#[test]
fn transfer_failure_poisons_slot_and_quarantines_without_leak() {
let fake = FakeTransfer::new();
fake.state.borrow_mut().fail_fill_layers = vec![7];
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
match db.prefetch(7, &layer(7)) {
Err(PrefillReject::Transfer(_)) => {}
other => panic!("expected Transfer error, got {other:?}"),
}
assert_eq!(db.metrics().poisoned, 1);
let ok = db.prefetch(1, &layer(1)).unwrap();
assert_ne!(ok.slot, 0, "poisoned slot 0 is not reused");
let _ = db.wait(&ok).unwrap();
db.release(ok).unwrap();
let ops = fake.ops();
assert!(ops.contains(&Op::Quarantine { slot: 0 }));
drop(db);
assert_no_leak(&fake);
}
#[test]
fn compute_wait_failure_poisons_the_ready_slot() {
let fake = FakeTransfer::new();
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
let t = db.prefetch(0, &layer(0)).unwrap();
fake.state.borrow_mut().fail_compute_wait = true;
match db.wait(&t) {
Err(PrefillReject::Transfer(_)) => {}
other => panic!("expected Transfer error, got {other:?}"),
}
assert_eq!(db.slot_status(t.slot), PrefillSlotStatus::Poisoned);
assert_eq!(db.metrics().poisoned, 1);
drop(db);
assert_no_leak(&fake);
}
#[test]
fn pool_capacity_decline_is_typed_and_makes_no_reservation() {
let fake = FakeTransfer::new();
fake.state.borrow_mut().pool_ok = false;
match PrefillDoubleBuffer::new(fake.clone(), 1 << 20) {
Err(PrefillReject::PoolCapacity { layer_bytes }) => assert_eq!(layer_bytes, 1 << 20),
other => panic!("expected PoolCapacity, got {other:?}"),
}
assert!(!fake.state.borrow().reserved);
assert_eq!(fake.ops(), Vec::new());
}
#[test]
fn reserve_transfer_failure_is_all_or_none() {
let fake = FakeTransfer::new();
fake.state.borrow_mut().reserve_fails = true;
match PrefillDoubleBuffer::new(fake.clone(), 1 << 20) {
Err(PrefillReject::Transfer(_)) => {}
other => panic!("expected Transfer error, got {other:?}"),
}
assert!(!fake.state.borrow().reserved);
}
#[test]
fn capture_active_rejects_reservation_and_prefetch() {
let fake = FakeTransfer::new();
fake.state.borrow_mut().capture = true;
match PrefillDoubleBuffer::new(fake.clone(), 1 << 20) {
Err(PrefillReject::CaptureActive) => {}
other => panic!("expected CaptureActive at reservation, got {other:?}"),
}
fake.state.borrow_mut().capture = false;
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
fake.state.borrow_mut().capture = true;
match db.prefetch(0, &layer(0)) {
Err(PrefillReject::CaptureActive) => {}
other => panic!("expected CaptureActive at prefetch, got {other:?}"),
}
assert_eq!(db.metrics().declined_capture, 1);
drop(db);
assert_no_leak(&fake);
}
#[test]
fn empty_layer_is_refused() {
let fake = FakeTransfer::new();
match PrefillDoubleBuffer::new(fake.clone(), 0) {
Err(PrefillReject::EmptyLayer) => {}
other => panic!("expected EmptyLayer at reservation, got {other:?}"),
}
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
match db.prefetch(0, &FakeLayer { id: 0, bytes: 0 }) {
Err(PrefillReject::EmptyLayer) => {}
other => panic!("expected EmptyLayer at prefetch, got {other:?}"),
}
assert_eq!(db.metrics().declined_empty, 1);
}
#[test]
fn depth_limit_refuses_a_third_unreleased_prefetch() {
let fake = FakeTransfer::new();
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
let _t0 = db.prefetch(0, &layer(0)).unwrap();
let _t1 = db.prefetch(1, &layer(1)).unwrap();
match db.prefetch(2, &layer(2)) {
Err(PrefillReject::SlotsBusy) => {}
other => panic!("expected SlotsBusy, got {other:?}"),
}
assert_eq!(db.metrics().declined_slots_busy, 1);
}
#[test]
fn instances_are_isolated() {
let fake_a = FakeTransfer::new();
let fake_b = FakeTransfer::new();
let mut a = PrefillDoubleBuffer::new(fake_a.clone(), 1 << 20).expect("reserve a");
let mut b = PrefillDoubleBuffer::new(fake_b.clone(), 1 << 20).expect("reserve b");
let ta = a.prefetch(10, &layer(10)).unwrap();
let tb = b.prefetch(20, &layer(20)).unwrap();
assert_eq!(a.wait(&ta).unwrap(), 10);
assert_eq!(b.wait(&tb).unwrap(), 20);
a.release(ta).unwrap();
b.release(tb).unwrap();
drop(a);
drop(b);
assert_no_leak(&fake_a);
assert_no_leak(&fake_b);
}
#[test]
fn release_fence_failure_poisons_slot() {
let fake = FakeTransfer::new();
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
let t = db.prefetch(0, &layer(0)).unwrap();
let _ = db.wait(&t).unwrap();
fake.state.borrow_mut().fail_release_fence = true;
match db.release(t) {
Err(PrefillReject::Transfer(_)) => {}
other => panic!("expected Transfer error, got {other:?}"),
}
assert_eq!(db.metrics().poisoned, 1);
drop(db);
assert_no_leak(&fake);
}
#[test]
fn reuse_wait_ns_is_reported_when_transfer_not_hidden() {
let fake = FakeTransfer::new();
fake.state.borrow_mut().reuse_wait_ns = 4_242;
let mut db = PrefillDoubleBuffer::new(fake.clone(), 1 << 20).expect("reserve");
cycle(&mut db, 0);
cycle(&mut db, 1);
cycle(&mut db, 2);
assert!(
db.metrics().reuse_wait_ns >= 4_242,
"unhidden reuse wait is surfaced"
);
drop(db);
assert_no_leak(&fake);
}
}