use super::reserved::{
ReservedTaskFuture, ReservedTaskJoined, ReservedTaskOutput, ReservedTaskTicket,
};
use saddle_admission::{
AdmissionError, ProfuseGwLightweightProcessOwner, StorageDemand, StoragePermit,
};
use std::alloc::Layout;
pub struct ReservedProcessTaskStorage {
permit: StoragePermit,
}
impl ReservedProcessTaskStorage {
pub fn pressure_sampler(&self, cpu_cores: usize) -> Result<crate::pressure::ManagedPressureSampler, AdmissionError> {
crate::pressure::ManagedPressureSampler::prepare(&self.permit, cpu_cores)
}
pub(crate) fn acquire(
process: &ProfuseGwLightweightProcessOwner,
) -> Result<Self, AdmissionError> {
Ok(Self {
permit: process.try_process_storage(StorageDemand::separate(&[])?)?,
})
}
pub fn normal<O: Send + 'static, T: Send + 'static>(
&self,
) -> Result<ReservedTaskCollection<O, T>, AdmissionError> {
if !cfg!(all(
target_arch = "x86_64",
target_os = "linux",
target_pointer_width = "64"
)) {
return Err(AdmissionError::InvalidConfiguration);
}
self.collection()
}
pub fn ingress<O: Send + 'static, T: Send + 'static>(
&self,
) -> Result<ReservedTaskCollection<O, T>, AdmissionError> {
self.collection()
}
pub fn rejection<O: Send + 'static, T: Send + 'static>(
&self,
) -> Result<ReservedTaskCollection<O, T>, AdmissionError> {
if !cfg!(all(
target_arch = "x86_64",
target_os = "linux",
target_pointer_width = "64"
)) {
return Err(AdmissionError::InvalidConfiguration);
}
self.collection()
}
pub fn ingress_queue<T>(&self) -> Result<ReservedIngressQueue<T>, AdmissionError> {
if std::mem::size_of::<T>() == 0 {
return Err(AdmissionError::InvalidConfiguration);
}
Ok(ReservedIngressQueue {
entries: std::collections::VecDeque::new(),
backing: None,
permit: self.permit.try_reserve(StorageDemand::separate(&[])?)?,
})
}
fn collection<O: Send + 'static, T: Send + 'static>(
&self,
) -> Result<ReservedTaskCollection<O, T>, AdmissionError> {
let permit = self.permit.try_reserve(StorageDemand::separate(&[(
Layout::from_size_align(72, 8).unwrap(),
1,
)])?)?;
Ok(ReservedTaskCollection {
tasks: tokio::task::JoinSet::new(),
tickets: Vec::new(),
backing: None,
permit,
})
}
}
pub struct ReservedIngressQueue<T> {
entries: std::collections::VecDeque<T>,
backing: Option<StoragePermit>,
permit: StoragePermit,
}
impl<T> ReservedIngressQueue<T> {
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn get_mut(&mut self, index: usize) -> Option<&mut T> {
self.entries.get_mut(index)
}
pub fn pop_front(&mut self) -> Option<T> {
self.entries.pop_front()
}
pub fn remove(&mut self, index: usize) -> Option<T> {
self.entries.remove(index)
}
pub fn push_back(&mut self, entry: T) -> Result<(), T> {
if self.entries.len() == self.entries.capacity() {
return Err(entry);
}
self.entries.push_back(entry);
Ok(())
}
pub fn push_front(&mut self, entry: T) -> Result<(), T> {
if self.entries.len() == self.entries.capacity() {
return Err(entry);
}
self.entries.push_front(entry);
Ok(())
}
pub fn try_capacity(&mut self, capacity: usize) -> Result<(), AdmissionError> {
if capacity <= self.entries.capacity() {
return Ok(());
}
let layout = Layout::array::<T>(capacity).map_err(|_| AdmissionError::SizeOverflow)?;
let backing = self
.permit
.try_reserve(StorageDemand::separate(&[(layout, 1)])?)?;
let mut replacement = std::collections::VecDeque::with_capacity(capacity);
assert_eq!(replacement.capacity(), capacity);
replacement.append(&mut self.entries);
drop(std::mem::replace(&mut self.entries, replacement));
drop(self.backing.replace(backing));
Ok(())
}
}
pub struct ReservedTaskCollection<O, T> {
tasks: tokio::task::JoinSet<ReservedTaskOutput<O, T>>,
tickets: Vec<(tokio::task::Id, ReservedTaskTicket<O>)>,
backing: Option<StoragePermit>,
permit: StoragePermit,
}
pub enum ReservedCollectionJoin<O, T> {
Matched(ReservedTaskJoined<O, T>),
ForeignTicket(
ReservedTaskTicket<O>,
Result<ReservedTaskOutput<O, T>, tokio::task::JoinError>,
),
UnknownTask(Result<ReservedTaskOutput<O, T>, tokio::task::JoinError>),
}
impl<O: Send + 'static, T: Send + 'static> ReservedTaskCollection<O, T> {
pub fn len(&self) -> usize {
self.tickets.len()
}
pub fn is_empty(&self) -> bool {
self.tickets.is_empty()
}
pub fn capacity(&self) -> usize {
self.tickets.capacity()
}
pub fn try_capacity(&mut self, capacity: usize) -> Result<(), AdmissionError> {
if capacity <= self.tickets.capacity() {
return Ok(());
}
let layout = Layout::array::<(tokio::task::Id, ReservedTaskTicket<O>)>(capacity)
.map_err(|_| AdmissionError::SizeOverflow)?;
let permit = self
.permit
.try_reserve(StorageDemand::separate(&[(layout, 1)])?)?;
let mut replacement = Vec::with_capacity(capacity);
assert_eq!(replacement.capacity(), capacity);
replacement.append(&mut self.tickets);
let old = std::mem::replace(&mut self.tickets, replacement);
drop(old);
drop(self.backing.replace(permit));
Ok(())
}
pub fn spawn(
&mut self,
future: ReservedTaskFuture<O, T>,
mut ticket: ReservedTaskTicket<O>,
) -> Result<
tokio::task::AbortHandle,
(
ReservedTaskFuture<O, T>,
ReservedTaskTicket<O>,
AdmissionError,
),
> {
if !ticket.matches_future(&future) {
return Err((future, ticket, AdmissionError::ForeignManagedSource));
}
if self.tickets.len() == self.tickets.capacity() {
let next = self
.tickets
.len()
.checked_add(1)
.ok_or(AdmissionError::SizeOverflow);
let result = next.and_then(|next| self.try_capacity(next));
if let Err(error) = result {
return Err((future, ticket, error));
}
}
let handle = self.tasks.spawn(future);
if ticket.bind(handle.id()).is_err() {
handle.abort();
}
self.tickets.push((handle.id(), ticket));
Ok(handle)
}
pub fn abort_all(&mut self) {
self.tasks.abort_all();
}
pub async fn join_next(&mut self) -> Option<ReservedCollectionJoin<O, T>> {
let result = self.tasks.join_next_with_id().await?;
let id = match &result {
Ok((id, _)) => *id,
Err(error) => error.id(),
};
let result = result.map(|(_, output)| output);
let Some(index) = self
.tickets
.iter()
.position(|(expected, _)| *expected == id)
else {
return Some(ReservedCollectionJoin::UnknownTask(result));
};
let (_, ticket) = self.tickets.swap_remove(index);
Some(match ticket.complete(result) {
Ok(joined) => ReservedCollectionJoin::Matched(joined),
Err((ticket, result)) => ReservedCollectionJoin::ForeignTicket(ticket, result),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::request_task::reserved::{
ReservedRequestFailure, tests::process, try_prepare_rejection_task,
};
#[test]
fn ingress_backing_grows_across_old_limits_and_retains_physical_charge() {
let process = process();
let storage = ReservedProcessTaskStorage::acquire(&process).unwrap();
let baseline = process.resource_snapshot().framework_charged;
let mut queue = storage.ingress_queue::<[u8; 32]>().unwrap();
queue.try_capacity(17).unwrap();
for value in 0..17 { queue.push_back([value; 32]).unwrap(); }
let before = process.resource_snapshot().framework_charged;
let next = StorageDemand::separate(&[(Layout::array::<[u8; 32]>(33).unwrap(), 1)]).unwrap();
let available = storage.permit.framework_budget_bytes() - before;
let filler_bytes = available - (next.bytes() - 1) - StoragePermit::layout().size();
let filler = storage.permit.try_reserve(StorageDemand::separate(&[
(Layout::array::<u8>(filler_bytes).unwrap(), 1),
]).unwrap()).unwrap();
let held = process.resource_snapshot().framework_charged;
assert!(matches!(queue.try_capacity(33), Err(AdmissionError::FrameworkReserveExceeded { .. })));
assert_eq!(queue.len(), 17);
assert_eq!(process.resource_snapshot().framework_charged, held);
drop(filler);
queue.try_capacity(33).unwrap();
assert_eq!(process.resource_snapshot().framework_charged, before + 16 * 32);
for value in 0..17 { assert_eq!(queue.pop_front(), Some([value; 32])); }
assert_eq!(process.resource_snapshot().framework_charged, before + 16 * 32);
drop(queue);
assert_eq!(process.resource_snapshot().framework_charged, baseline);
drop(storage);
process.finish().unwrap();
}
#[test]
fn real_rejection_collection_retains_capacity_until_join_and_storage_drop() {
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap()
.block_on(async {
let process = process();
let storage = ReservedProcessTaskStorage::acquire(&process).unwrap();
let mut tasks = storage.rejection::<(), u32>().unwrap();
tasks.try_capacity(1).unwrap();
tasks.try_capacity(2).unwrap();
assert_eq!(tasks.capacity(), 2);
tasks.try_capacity(17).unwrap();
tasks.try_capacity(33).unwrap();
assert_eq!(tasks.capacity(), 33, "old slot boundary is not a resource limit");
let impossible = tasks.permit.framework_budget_bytes();
assert!(tasks.try_capacity(impossible).is_err());
assert_eq!(tasks.capacity(), 33, "failed growth preserves old backing");
let (root, future, ticket) = try_prepare_rejection_task(
&process.verified_profile(),
saddle_core::request_context::ContextLabel::checked("app").unwrap(),
None,
Layout::new::<std::future::Ready<Result<u32, ReservedRequestFailure>>>(),
&[],
(),
|_, _| Box::pin(std::future::ready(Ok(41))),
)
.ok()
.unwrap();
tasks.spawn(future, ticket).ok().unwrap();
tokio::task::yield_now().await;
assert_eq!(tasks.len(), 1, "completed is not joined");
let Some(ReservedCollectionJoin::Matched(joined)) = tasks.join_next().await else {
panic!("matching ticket")
};
drop(root);
let recovered = joined.recover(Default::default()).ok().unwrap();
assert_eq!(recovered.result.unwrap().ok().unwrap(), 41);
assert_eq!(tasks.len(), 0);
assert_eq!(tasks.capacity(), 33, "remove does not refund backing");
drop((tasks, storage));
process.finish().unwrap();
});
}
}