use std::ops::{Deref, DerefMut};
use crate::runtime::session::RuntimeSession;
use crate::runtime::tensor::Tensor;
pub struct SessionTriplet {
pub(crate) encoder: Box<dyn RuntimeSession>,
pub(crate) decoder: Option<Box<dyn RuntimeSession>>,
pub(crate) joiner: Option<Box<dyn RuntimeSession>>,
pub(crate) encoder_inputs: Vec<Tensor>,
}
#[derive(Debug)]
pub enum PoolError {
Closed,
}
impl std::fmt::Display for PoolError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PoolError::Closed => write!(f, "session pool is closed"),
}
}
}
impl std::error::Error for PoolError {}
pub struct Pool<T> {
inner: std::sync::Arc<PoolInner<T>>,
}
struct PoolInner<T> {
items: parking_lot::Mutex<std::collections::VecDeque<T>>,
waiters: parking_lot::Mutex<std::collections::VecDeque<Waiter<T>>>,
closed: std::sync::atomic::AtomicBool,
total: usize,
}
enum Waiter<T> {
#[cfg(feature = "async-pool")]
Async(tokio::sync::oneshot::Sender<T>),
Blocking(std::sync::mpsc::Sender<T>),
}
pub type SessionPool = Pool<SessionTriplet>;
impl<T: Send> Pool<T> {
pub fn new(items: Vec<T>) -> Self {
let total = items.len();
Self {
inner: std::sync::Arc::new(PoolInner {
items: parking_lot::Mutex::new(std::collections::VecDeque::from(items)),
waiters: parking_lot::Mutex::new(std::collections::VecDeque::new()),
closed: std::sync::atomic::AtomicBool::new(false),
total,
}),
}
}
#[cfg(feature = "async-pool")]
pub async fn checkout(&self) -> Result<PoolGuard<T>, PoolError> {
{
let mut items = self.inner.items.lock();
if self.inner.closed.load(std::sync::atomic::Ordering::SeqCst) {
return Err(PoolError::Closed);
}
if let Some(item) = items.pop_front() {
return Ok(PoolGuard::new(self.inner.clone(), item));
}
}
let (tx, rx) = tokio::sync::oneshot::channel();
{
let mut waiters = self.inner.waiters.lock();
if self.inner.closed.load(std::sync::atomic::Ordering::SeqCst) {
return Err(PoolError::Closed);
}
let mut items = self.inner.items.lock();
if let Some(item) = items.pop_front() {
drop(items);
drop(waiters);
return Ok(PoolGuard::new(self.inner.clone(), item));
}
waiters.push_back(Waiter::Async(tx));
}
match rx.await {
Ok(item) => Ok(PoolGuard::new(self.inner.clone(), item)),
Err(_) => Err(PoolError::Closed),
}
}
pub fn checkout_blocking(&self) -> Result<PoolGuard<T>, PoolError> {
{
let mut items = self.inner.items.lock();
if self.inner.closed.load(std::sync::atomic::Ordering::SeqCst) {
return Err(PoolError::Closed);
}
if let Some(item) = items.pop_front() {
return Ok(PoolGuard::new(self.inner.clone(), item));
}
}
let (tx, rx) = std::sync::mpsc::channel();
{
let mut waiters = self.inner.waiters.lock();
if self.inner.closed.load(std::sync::atomic::Ordering::SeqCst) {
return Err(PoolError::Closed);
}
let mut items = self.inner.items.lock();
if let Some(item) = items.pop_front() {
drop(items);
drop(waiters);
return Ok(PoolGuard::new(self.inner.clone(), item));
}
waiters.push_back(Waiter::Blocking(tx));
}
match rx.recv() {
Ok(item) => Ok(PoolGuard::new(self.inner.clone(), item)),
Err(_) => Err(PoolError::Closed),
}
}
pub fn close(&self) {
self.inner
.closed
.store(true, std::sync::atomic::Ordering::SeqCst);
let mut waiters = self.inner.waiters.lock();
waiters.clear();
}
pub fn total(&self) -> usize {
self.inner.total
}
pub fn available(&self) -> usize {
let items = self.inner.items.lock();
items.len()
}
pub fn waiters(&self) -> usize {
let waiters = self.inner.waiters.lock();
waiters.len()
}
}
impl<T> PoolInner<T> {
fn checkin(&self, mut item: T) {
if self.closed.load(std::sync::atomic::Ordering::SeqCst) {
return;
}
loop {
let mut waiters = self.waiters.lock();
if let Some(waiter) = waiters.pop_front() {
drop(waiters);
match waiter {
#[cfg(feature = "async-pool")]
Waiter::Async(tx) => {
if let Err(returned_item) = tx.send(item) {
item = returned_item;
continue;
}
}
Waiter::Blocking(tx) => {
if let Err(std::sync::mpsc::SendError(returned_item)) = tx.send(item) {
item = returned_item;
continue;
}
}
}
} else {
drop(waiters);
let mut items = self.items.lock();
items.push_back(item);
}
break;
}
}
}
pub struct PoolGuard<T> {
inner: Option<std::sync::Arc<PoolInner<T>>>,
item: Option<T>,
}
impl<T> PoolGuard<T> {
fn new(inner: std::sync::Arc<PoolInner<T>>, item: T) -> Self {
Self {
inner: Some(inner),
item: Some(item),
}
}
pub fn into_owned(mut self) -> OwnedReservation<T> {
let item = self
.item
.take()
.unwrap_or_else(|| unreachable!("PoolGuard::into_owned called after drop"));
let inner = self.inner.take().unwrap();
OwnedReservation {
inner,
item: Some(item),
}
}
}
impl<T> Deref for PoolGuard<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
self.item
.as_ref()
.unwrap_or_else(|| unreachable!("PoolGuard accessed after item taken"))
}
}
impl<T> DerefMut for PoolGuard<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
self.item
.as_mut()
.unwrap_or_else(|| unreachable!("PoolGuard accessed after item taken"))
}
}
impl<T> Drop for PoolGuard<T> {
fn drop(&mut self) {
if let (Some(inner), Some(item)) = (self.inner.take(), self.item.take()) {
inner.checkin(item);
}
}
}
pub struct OwnedReservation<T> {
inner: std::sync::Arc<PoolInner<T>>,
item: Option<T>,
}
impl<T> std::ops::Deref for OwnedReservation<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
self.item
.as_ref()
.unwrap_or_else(|| unreachable!("OwnedReservation accessed after checkin"))
}
}
impl<T> std::ops::DerefMut for OwnedReservation<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
self.item
.as_mut()
.unwrap_or_else(|| unreachable!("OwnedReservation accessed after checkin"))
}
}
impl<T> OwnedReservation<T> {
pub fn checkin(mut self) {
if let Some(item) = self.item.take() {
self.inner.checkin(item);
}
}
}
impl<T> Drop for OwnedReservation<T> {
fn drop(&mut self) {
if let Some(item) = self.item.take() {
self.inner.checkin(item);
}
}
}