use super::reserved::{
ReservedTaskFuture, ReservedTaskJoined, ReservedTaskOutput, ReservedTaskTicket,
};
use saddle_admission::{
AdmissionError, ProfuseGwLightweightProcessOwner, StorageDemand, StoragePermit,
};
use std::alloc::Layout;
pub struct ReservedProcessTaskStorage {
permit: StoragePermit,
normal_limit: usize,
}
impl ReservedProcessTaskStorage {
pub(crate) fn acquire(
process: &ProfuseGwLightweightProcessOwner,
) -> Result<Self, AdmissionError> {
Ok(Self {
permit: process.try_process_storage(StorageDemand::separate(&[])?)?,
normal_limit: process.capacity_snapshot().active_limit(),
})
}
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(self.normal_limit)
}
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(16)
}
fn collection<O: Send + 'static, T: Send + 'static>(
&self,
limit: usize,
) -> 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,
limit,
})
}
}
pub struct ReservedTaskCollection<O, T> {
tasks: tokio::task::JoinSet<ReservedTaskOutput<O, T>>,
tickets: Vec<(tokio::task::Id, ReservedTaskTicket<O>)>,
backing: Option<StoragePermit>,
permit: StoragePermit,
limit: usize,
}
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(());
}
if capacity > self.limit {
return Err(AdmissionError::CapacityRejected);
}
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.limit {
return Err((future, ticket, AdmissionError::CapacityRejected));
}
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 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);
assert!(tasks.try_capacity(17).is_err());
assert_eq!(tasks.capacity(), 2, "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(), 2, "remove does not refund backing");
drop((tasks, storage));
process.finish().unwrap();
});
}
}