use std::cell::{Cell, RefCell};
use std::future::Future;
use std::pin::Pin;
use std::rc::Rc;
use std::task::{Context, Poll};
use crate::event::{new_event, EventAwaitable, EventTrigger};
use super::wait_queue::WaitQueue;
struct Holder {
priority: u32,
id: u64,
trigger: Option<EventTrigger>,
preempted: Rc<Cell<bool>>,
}
struct PreemptiveState {
wq: WaitQueue<u32>,
holders: Vec<Holder>,
next_holder_id: u64,
}
impl PreemptiveState {
fn victim_index(&self, incoming: u32) -> Option<usize> {
self.holders
.iter()
.enumerate()
.filter(|(_, h)| h.priority > incoming)
.max_by_key(|(_, h)| h.priority)
.map(|(i, _)| i)
}
}
#[derive(Clone)]
pub struct PreemptiveResource {
state: Rc<RefCell<PreemptiveState>>,
}
impl std::fmt::Debug for PreemptiveResource {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut d = f.debug_struct("PreemptiveResource");
if let Ok(s) = self.state.try_borrow() {
d.field("in_use", &s.wq.in_use())
.field("capacity", &s.wq.capacity())
.field("queue_len", &s.wq.live_waiters());
}
d.finish_non_exhaustive()
}
}
impl PreemptiveResource {
#[must_use]
pub fn new(capacity: usize) -> Self {
assert!(
capacity > 0,
"PreemptiveResource capacity must be at least 1"
);
PreemptiveResource {
state: Rc::new(RefCell::new(PreemptiveState {
wq: WaitQueue::new(capacity),
holders: Vec::new(),
next_holder_id: 0,
})),
}
}
#[must_use = "futures do nothing unless awaited"]
pub fn request(&self, priority: u32) -> PreemptiveRequest {
PreemptiveRequest {
state: Rc::clone(&self.state),
priority,
registered: false,
consumed: false,
canceled: Rc::new(Cell::new(false)),
granted: Rc::new(Cell::new(false)),
}
}
#[must_use]
pub fn in_use(&self) -> usize {
self.state.borrow().wq.in_use()
}
#[must_use]
pub fn capacity(&self) -> usize {
self.state.borrow().wq.capacity()
}
#[must_use]
pub fn queue_len(&self) -> usize {
self.state.borrow().wq.live_waiters()
}
}
fn grant(state_rc: &Rc<RefCell<PreemptiveState>>, priority: u32) -> PreemptiveGuard {
let (trigger, awaitable) = new_event();
let preempted = Rc::new(Cell::new(false));
let id = {
let mut state = state_rc.borrow_mut();
let id = state.next_holder_id;
state.next_holder_id += 1;
state.holders.push(Holder {
priority,
id,
trigger: Some(trigger),
preempted: Rc::clone(&preempted),
});
id
};
PreemptiveGuard {
state: Rc::clone(state_rc),
id,
signal: awaitable,
preempted,
}
}
pub struct PreemptiveRequest {
state: Rc<RefCell<PreemptiveState>>,
priority: u32,
registered: bool,
consumed: bool,
canceled: Rc<Cell<bool>>,
granted: Rc<Cell<bool>>,
}
impl std::fmt::Debug for PreemptiveRequest {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PreemptiveRequest")
.field("priority", &self.priority)
.field("registered", &self.registered)
.field("granted", &self.granted.get())
.finish_non_exhaustive()
}
}
impl Future for PreemptiveRequest {
type Output = PreemptiveGuard;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<PreemptiveGuard> {
if self.granted.get() {
self.consumed = true;
return Poll::Ready(grant(&self.state, self.priority));
}
enum Outcome {
Free,
Preempt(usize),
Block,
}
let outcome = {
let mut state = self.state.borrow_mut();
if state.wq.try_acquire() {
Outcome::Free
} else if let Some(idx) = state.victim_index(self.priority) {
Outcome::Preempt(idx)
} else {
if !self.registered {
state.wq.register(
self.priority,
cx.waker().clone(),
Rc::clone(&self.canceled),
Rc::clone(&self.granted),
);
}
Outcome::Block
}
};
match outcome {
Outcome::Free => {
self.consumed = true;
Poll::Ready(grant(&self.state, self.priority))
}
Outcome::Preempt(idx) => {
let victim = {
let mut state = self.state.borrow_mut();
state.holders.remove(idx)
};
victim.preempted.set(true);
if let Some(trigger) = victim.trigger {
trigger.fire();
}
self.consumed = true;
Poll::Ready(grant(&self.state, self.priority))
}
Outcome::Block => {
self.registered = true;
Poll::Pending
}
}
}
}
impl Drop for PreemptiveRequest {
fn drop(&mut self) {
if self.consumed {
return; }
if self.granted.get() {
self.state.borrow_mut().wq.release();
} else if self.registered {
self.canceled.set(true);
}
}
}
pub struct PreemptiveGuard {
state: Rc<RefCell<PreemptiveState>>,
id: u64,
signal: EventAwaitable,
preempted: Rc<Cell<bool>>,
}
impl std::fmt::Debug for PreemptiveGuard {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PreemptiveGuard")
.field("id", &self.id)
.field("preempted", &self.preempted.get())
.finish_non_exhaustive()
}
}
impl PreemptiveGuard {
#[must_use = "futures do nothing unless awaited"]
pub fn preempted(&self) -> EventAwaitable {
self.signal.clone()
}
#[must_use]
pub fn is_preempted(&self) -> bool {
self.preempted.get()
}
}
impl Drop for PreemptiveGuard {
fn drop(&mut self) {
let mut state = self.state.borrow_mut();
if self.preempted.get() {
return;
}
if let Some(pos) = state.holders.iter().position(|h| h.id == self.id) {
state.holders.remove(pos);
}
state.wq.release();
}
}