use std::{
collections::{BTreeMap, HashMap, VecDeque},
sync::{Arc, Mutex},
};
use tokio::sync::Notify;
use crate::{Error, Frame, StreamId};
struct StreamSlot {
priority: u8,
frames: VecDeque<Frame>,
}
struct Inner {
bands: BTreeMap<u8, VecDeque<StreamId>>,
streams: HashMap<StreamId, StreamSlot>,
len: usize,
closed: bool,
}
impl Inner {
fn arm(&mut self, id: StreamId, band: u8) {
self.bands.entry(band).or_default().push_back(id);
}
}
#[derive(Clone)]
pub struct PriorityQueue {
inner: Arc<Mutex<Inner>>,
non_empty: Arc<Notify>,
has_space: Arc<Notify>,
capacity: usize,
}
impl PriorityQueue {
pub fn new(capacity: usize) -> Self {
Self {
inner: Arc::new(Mutex::new(Inner {
bands: BTreeMap::new(),
streams: HashMap::new(),
len: 0,
closed: false,
})),
non_empty: Arc::new(Notify::new()),
has_space: Arc::new(Notify::new()),
capacity,
}
}
pub async fn reserve(&self) -> Result<Permit, Error> {
loop {
let notified = self.has_space.notified();
{
let mut inner = self.inner.lock().unwrap();
if inner.closed {
return Err(Error::Closed);
}
if inner.len < self.capacity {
inner.len += 1;
return Ok(Permit {
queue: self.clone(),
armed: true,
});
}
}
notified.await;
}
}
pub fn push_now(&self, priority: u8, id: StreamId, frame: Frame) -> Result<(), Error> {
let mut inner = self.inner.lock().unwrap();
if inner.closed {
return Err(Error::Closed);
}
self.push_locked(&mut inner, priority, id, frame);
Ok(())
}
fn push_locked(&self, inner: &mut Inner, priority: u8, id: StreamId, frame: Frame) {
self.enqueue_locked(inner, priority, id, frame);
inner.len += 1;
}
fn enqueue_locked(&self, inner: &mut Inner, priority: u8, id: StreamId, frame: Frame) {
match inner.streams.get_mut(&id) {
Some(slot) => {
slot.frames.push_back(frame);
}
None => {
let mut frames = VecDeque::new();
frames.push_back(frame);
inner.streams.insert(id, StreamSlot { priority, frames });
inner.arm(id, priority);
}
}
self.non_empty.notify_one();
}
pub async fn pop(&self) -> Option<Frame> {
loop {
let notified = self.non_empty.notified();
{
let mut inner = self.inner.lock().unwrap();
if let Some(frame) = self.pop_locked(&mut inner) {
return Some(frame);
}
if inner.closed {
return None;
}
}
notified.await;
}
}
fn pop_locked(&self, inner: &mut Inner) -> Option<Frame> {
let (&band, queue) = inner.bands.iter_mut().next_back()?;
let id = queue.pop_front().expect("scheduled band must be non-empty");
if queue.is_empty() {
inner.bands.remove(&band);
}
let slot = inner
.streams
.get_mut(&id)
.expect("scheduled stream must have a slot");
let frame = slot
.frames
.pop_front()
.expect("scheduled slot must be non-empty");
if slot.frames.is_empty() {
inner.streams.remove(&id);
} else {
let priority = slot.priority;
inner.arm(id, priority);
}
inner.len -= 1;
self.has_space.notify_one();
Some(frame)
}
pub fn set_priority(&self, id: StreamId, new: u8) {
let mut inner = self.inner.lock().unwrap();
let old = match inner.streams.get(&id) {
Some(slot) => slot.priority,
None => return,
};
if old == new {
return;
}
if let Some(queue) = inner.bands.get_mut(&old) {
if let Some(pos) = queue.iter().position(|&s| s == id) {
queue.remove(pos);
if queue.is_empty() {
inner.bands.remove(&old);
}
inner.bands.entry(new).or_default().push_back(id);
}
}
if let Some(slot) = inner.streams.get_mut(&id) {
slot.priority = new;
}
}
pub fn remove(&self, id: StreamId) -> u64 {
let mut inner = self.inner.lock().unwrap();
let Some(slot) = inner.streams.remove(&id) else {
return 0;
};
let removed = slot.frames.len();
let removed_bytes = slot
.frames
.iter()
.map(|frame| match frame {
Frame::Stream(stream) => stream.data.len() as u64,
_ => 0,
})
.sum();
if let Some(queue) = inner.bands.get_mut(&slot.priority) {
if let Some(pos) = queue.iter().position(|&s| s == id) {
queue.remove(pos);
}
if queue.is_empty() {
inner.bands.remove(&slot.priority);
}
}
inner.len -= removed;
drop(inner);
for _ in 0..removed {
self.has_space.notify_one();
}
removed_bytes
}
pub fn close(&self) {
{
let mut inner = self.inner.lock().unwrap();
inner.closed = true;
}
self.non_empty.notify_waiters();
self.has_space.notify_waiters();
}
}
pub struct Permit {
queue: PriorityQueue,
armed: bool,
}
impl Permit {
pub fn send(mut self, priority: u8, id: StreamId, frame: Frame) -> Result<(), Error> {
let mut inner = self.queue.inner.lock().unwrap();
if inner.closed {
drop(inner);
return Err(Error::Closed);
}
self.armed = false;
self.queue.enqueue_locked(&mut inner, priority, id, frame);
Ok(())
}
}
impl Drop for Permit {
fn drop(&mut self) {
if !self.armed {
return;
}
{
let mut inner = self.queue.inner.lock().unwrap();
inner.len -= 1;
}
self.queue.has_space.notify_one();
}
}
#[cfg(test)]
mod tests {
use super::*;
use bytes::Bytes;
use crate::proto::Stream;
use crate::{StreamDir, StreamId};
fn sid(index: u64) -> StreamId {
StreamId::new(index, StreamDir::Uni, false)
}
fn frame(id: StreamId, tag: u8) -> Frame {
Frame::Stream(Stream {
id,
offset: 0,
data: Bytes::copy_from_slice(&[tag]),
fin: false,
})
}
fn tag_of(frame: &Frame) -> u8 {
match frame {
Frame::Stream(s) => s.data[0],
_ => panic!("expected stream frame"),
}
}
fn id_of(frame: &Frame) -> StreamId {
match frame {
Frame::Stream(s) => s.id,
_ => panic!("expected stream frame"),
}
}
async fn push(
q: &PriorityQueue,
priority: u8,
id: StreamId,
frame: Frame,
) -> Result<(), Error> {
q.reserve().await?.send(priority, id, frame)
}
#[tokio::test]
async fn higher_priority_first() {
let q = PriorityQueue::new(8);
let lo = sid(0);
let hi = sid(1);
push(&q, 10, lo, frame(lo, b'l')).await.unwrap();
push(&q, 200, hi, frame(hi, b'h')).await.unwrap();
assert_eq!(tag_of(&q.pop().await.unwrap()), b'h');
assert_eq!(tag_of(&q.pop().await.unwrap()), b'l');
}
#[tokio::test]
async fn equal_priority_round_robin() {
let q = PriorityQueue::new(8);
let a = sid(0);
let b = sid(1);
push(&q, 5, a, frame(a, 1)).await.unwrap();
push(&q, 5, a, frame(a, 2)).await.unwrap();
push(&q, 5, b, frame(b, 1)).await.unwrap();
push(&q, 5, b, frame(b, 2)).await.unwrap();
assert_eq!(id_of(&q.pop().await.unwrap()), a);
assert_eq!(id_of(&q.pop().await.unwrap()), b);
assert_eq!(id_of(&q.pop().await.unwrap()), a);
assert_eq!(id_of(&q.pop().await.unwrap()), b);
}
#[tokio::test]
async fn per_stream_fifo_preserved() {
let q = PriorityQueue::new(8);
let a = sid(0);
for i in 0..4u8 {
push(&q, 5, a, frame(a, i)).await.unwrap();
}
for i in 0..4u8 {
assert_eq!(tag_of(&q.pop().await.unwrap()), i);
}
}
#[tokio::test]
async fn set_priority_moves_pointer_not_frames() {
let q = PriorityQueue::new(8);
let lo = sid(0);
let hi = sid(1);
push(&q, 10, lo, frame(lo, 1)).await.unwrap();
push(&q, 10, lo, frame(lo, 2)).await.unwrap();
push(&q, 20, hi, frame(hi, 1)).await.unwrap();
q.set_priority(lo, 100);
assert_eq!(id_of(&q.pop().await.unwrap()), lo);
assert_eq!(tag_of(&q.pop().await.unwrap()), 2); assert_eq!(id_of(&q.pop().await.unwrap()), hi);
}
#[tokio::test]
async fn set_priority_unknown_stream_is_noop() {
let q = PriorityQueue::new(8);
q.set_priority(sid(99), 50); }
#[tokio::test]
async fn close_unblocks_pop() {
let q = PriorityQueue::new(8);
let q2 = q.clone();
let handle = tokio::spawn(async move { q2.pop().await });
tokio::task::yield_now().await;
q.close();
assert!(handle.await.unwrap().is_none());
}
#[tokio::test]
async fn close_unblocks_push() {
let q = PriorityQueue::new(1);
let a = sid(0);
push(&q, 5, a, frame(a, 1)).await.unwrap();
let q2 = q.clone();
let handle = tokio::spawn(async move { push(&q2, 5, sid(1), frame(sid(1), 2)).await });
tokio::task::yield_now().await;
q.close();
assert!(matches!(handle.await.unwrap(), Err(Error::Closed)));
}
#[tokio::test]
async fn backpressure_blocks_at_capacity() {
let q = PriorityQueue::new(2);
let a = sid(0);
push(&q, 5, a, frame(a, 1)).await.unwrap();
push(&q, 5, a, frame(a, 2)).await.unwrap();
let q2 = q.clone();
let pushing = tokio::spawn(async move { push(&q2, 5, a, frame(a, 3)).await });
tokio::task::yield_now().await;
assert!(!pushing.is_finished(), "push should block while full");
q.pop().await.unwrap();
pushing.await.unwrap().unwrap();
}
#[tokio::test]
async fn remove_drops_a_streams_queued_frames() {
let q = PriorityQueue::new(8);
let a = sid(0);
let b = sid(1);
push(&q, 5, a, frame(a, 1)).await.unwrap();
push(&q, 5, a, frame(a, 2)).await.unwrap();
push(&q, 5, b, frame(b, 9)).await.unwrap();
assert_eq!(q.remove(a), 2);
assert_eq!(id_of(&q.pop().await.unwrap()), b);
q.close();
assert!(q.pop().await.is_none());
}
#[tokio::test]
async fn remove_frees_capacity_for_blocked_producers() {
let q = PriorityQueue::new(2);
let a = sid(0);
let b = sid(1);
push(&q, 5, a, frame(a, 1)).await.unwrap();
push(&q, 5, a, frame(a, 2)).await.unwrap();
let q2 = q.clone();
let pushing = tokio::spawn(async move { push(&q2, 5, b, frame(b, 1)).await });
tokio::task::yield_now().await;
assert!(!pushing.is_finished(), "push should block while full");
assert_eq!(q.remove(a), 2);
pushing.await.unwrap().unwrap();
assert_eq!(id_of(&q.pop().await.unwrap()), b);
}
#[tokio::test]
async fn remove_unknown_stream_is_noop() {
let q = PriorityQueue::new(8);
assert_eq!(q.remove(sid(99)), 0); }
#[tokio::test]
async fn reserved_permit_holds_capacity() {
let q = PriorityQueue::new(1);
let a = sid(0);
let permit = q.reserve().await.unwrap();
let q2 = q.clone();
let blocked = tokio::spawn(async move { q2.reserve().await.map(|_| ()) });
tokio::task::yield_now().await;
assert!(!blocked.is_finished(), "reserve should block while full");
permit.send(5, a, frame(a, 1)).unwrap();
tokio::task::yield_now().await;
assert!(
!blocked.is_finished(),
"sending a permit must not free its slot"
);
q.pop().await.unwrap();
blocked.await.unwrap().unwrap();
}
#[tokio::test]
async fn dropped_permit_returns_capacity() {
let q = PriorityQueue::new(1);
let q2 = q.clone();
let blocked = tokio::spawn(async move { q2.reserve().await.map(|_| ()) });
{
let _permit = q.reserve().await.unwrap();
tokio::task::yield_now().await;
assert!(!blocked.is_finished(), "reserve should block while full");
}
blocked.await.unwrap().unwrap();
}
#[tokio::test]
async fn cancelled_reserve_holds_nothing() {
let q = PriorityQueue::new(1);
let a = sid(0);
push(&q, 5, a, frame(a, 1)).await.unwrap();
tokio::select! {
_ = q.reserve() => panic!("reserve should not succeed while full"),
_ = tokio::task::yield_now() => {}
}
q.pop().await.unwrap();
push(&q, 5, a, frame(a, 2)).await.unwrap();
assert_eq!(tag_of(&q.pop().await.unwrap()), 2);
}
#[tokio::test]
async fn send_after_close_fails() {
let q = PriorityQueue::new(4);
let a = sid(0);
let permit = q.reserve().await.unwrap();
q.close();
assert!(q.pop().await.is_none());
assert!(matches!(permit.send(5, a, frame(a, 1)), Err(Error::Closed)));
assert!(
q.pop().await.is_none(),
"a rejected send must queue nothing"
);
}
#[tokio::test]
async fn drop_after_close_is_clean() {
let q = PriorityQueue::new(4);
let permit = q.reserve().await.unwrap();
q.close();
drop(permit);
assert!(q.pop().await.is_none());
}
#[tokio::test]
async fn close_unblocks_reserve() {
let q = PriorityQueue::new(1);
let a = sid(0);
push(&q, 5, a, frame(a, 1)).await.unwrap();
let q2 = q.clone();
let handle = tokio::spawn(async move { q2.reserve().await.map(|_| ()) });
tokio::task::yield_now().await;
q.close();
assert!(matches!(handle.await.unwrap(), Err(Error::Closed)));
}
}