use crate::agent::{AbortSignal, AgentMessage, GetQueuedMessagesFn, QueueKind, QueueMode};
use parking_lot::Mutex;
use std::collections::VecDeque;
use std::sync::atomic::{AtomicU64, Ordering};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct QueuedMessageId(u64);
impl QueuedMessageId {
pub fn as_u64(self) -> u64 {
self.0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct QueuedMessageHandle {
pub kind: QueueKind,
pub id: QueuedMessageId,
}
#[derive(Debug, Clone)]
struct QueuedMessageItem {
id: QueuedMessageId,
message: AgentMessage,
}
pub trait DrainStrategy: Send + Sync {
fn select(&self, buffer: &mut VecDeque<AgentMessage>) -> Vec<AgentMessage>;
fn name(&self) -> &'static str;
}
#[allow(dead_code)]
pub struct DrainAll;
impl DrainStrategy for DrainAll {
fn select(&self, buffer: &mut VecDeque<AgentMessage>) -> Vec<AgentMessage> {
buffer.drain(..).collect()
}
fn name(&self) -> &'static str {
"all"
}
}
#[allow(dead_code)]
pub struct DrainOne;
impl DrainStrategy for DrainOne {
fn select(&self, buffer: &mut VecDeque<AgentMessage>) -> Vec<AgentMessage> {
buffer.pop_front().into_iter().collect()
}
fn name(&self) -> &'static str {
"one_at_a_time"
}
}
#[allow(dead_code)]
pub struct DrainBatch {
pub max: usize,
}
impl DrainStrategy for DrainBatch {
fn select(&self, buffer: &mut VecDeque<AgentMessage>) -> Vec<AgentMessage> {
let take = self.max.min(buffer.len());
buffer.drain(..take).collect()
}
fn name(&self) -> &'static str {
"batch"
}
}
impl QueueMode {
#[allow(dead_code)]
pub(crate) fn to_strategy(self) -> Box<dyn DrainStrategy> {
match self {
QueueMode::All => Box::new(DrainAll),
QueueMode::OneAtATime => Box::new(DrainOne),
}
}
}
#[derive(Debug, Clone)]
pub struct BackpressureConfig {
pub max_depth: usize,
pub overflow: OverflowBehavior,
}
impl Default for BackpressureConfig {
fn default() -> Self {
Self {
max_depth: 0,
overflow: OverflowBehavior::Unlimited,
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub enum OverflowBehavior {
#[default]
Unlimited,
DropOldest,
Reject,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct QueueFullError {
pub current_depth: usize,
pub max_depth: usize,
}
impl std::fmt::Display for QueueFullError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"queue full: depth {} >= max {}",
self.current_depth, self.max_depth
)
}
}
impl std::error::Error for QueueFullError {}
pub(crate) struct MessageQueue {
buffer: Mutex<VecDeque<QueuedMessageItem>>,
kind: QueueKind,
backpressure: Mutex<BackpressureConfig>,
next_id: AtomicU64,
}
impl MessageQueue {
pub fn new(kind: QueueKind) -> Self {
Self {
buffer: Mutex::new(VecDeque::new()),
kind,
backpressure: Mutex::new(BackpressureConfig::default()),
next_id: AtomicU64::new(1),
}
}
fn next_message_id(&self) -> QueuedMessageId {
QueuedMessageId(self.next_id.fetch_add(1, Ordering::Relaxed))
}
fn make_item(&self, message: AgentMessage) -> QueuedMessageItem {
QueuedMessageItem {
id: self.next_message_id(),
message,
}
}
#[allow(dead_code)]
pub fn kind(&self) -> QueueKind {
self.kind
}
pub fn push(&self, message: AgentMessage) -> QueuedMessageId {
let item = self.make_item(message);
let id = item.id;
self.buffer.lock().push_back(item);
id
}
#[allow(dead_code)]
pub fn push_many(&self, messages: impl IntoIterator<Item = AgentMessage>) {
let mut buf = self.buffer.lock();
buf.extend(messages.into_iter().map(|message| self.make_item(message)));
}
pub fn try_push(&self, message: AgentMessage) -> Result<QueuedMessageId, QueueFullError> {
let bp = self.backpressure.lock().clone();
let item = self.make_item(message);
let id = item.id;
if bp.max_depth == 0 || bp.overflow == OverflowBehavior::Unlimited {
self.buffer.lock().push_back(item);
return Ok(id);
}
let mut buf = self.buffer.lock();
if buf.len() >= bp.max_depth {
match bp.overflow {
OverflowBehavior::Reject => {
return Err(QueueFullError {
current_depth: buf.len(),
max_depth: bp.max_depth,
});
}
OverflowBehavior::DropOldest => {
buf.pop_front();
}
OverflowBehavior::Unlimited => unreachable!(),
}
}
buf.push_back(item);
Ok(id)
}
pub fn set_backpressure(&self, config: BackpressureConfig) {
*self.backpressure.lock() = config;
}
#[allow(dead_code)]
pub fn backpressure(&self) -> BackpressureConfig {
self.backpressure.lock().clone()
}
#[allow(dead_code)]
pub fn drain_with_strategy(&self, strategy: &dyn DrainStrategy) -> Vec<AgentMessage> {
let mut buf = self.buffer.lock();
let mut messages: VecDeque<AgentMessage> = buf.drain(..).map(|item| item.message).collect();
let selected = strategy.select(&mut messages);
buf.extend(messages.into_iter().map(|message| self.make_item(message)));
selected
}
pub fn drain_local(&self, mode: QueueMode) -> Vec<AgentMessage> {
let mut buf = self.buffer.lock();
match mode {
QueueMode::All => buf.drain(..).map(|item| item.message).collect(),
QueueMode::OneAtATime => {
if let Some(first) = buf.pop_front() {
vec![first.message]
} else {
Vec::new()
}
}
}
}
pub async fn drain(
&self,
mode: QueueMode,
supplier: &Option<GetQueuedMessagesFn>,
abort: AbortSignal,
) -> Vec<AgentMessage> {
let local = self.drain_local(mode);
let dynamic = match supplier {
Some(s) if mode == QueueMode::All || local.is_empty() => s(abort).await,
_ => Vec::new(),
};
match mode {
QueueMode::All => {
let mut merged = local;
merged.extend(dynamic);
merged
}
QueueMode::OneAtATime => {
if !local.is_empty() {
local
} else {
dynamic.into_iter().take(1).collect()
}
}
}
}
pub fn is_empty(&self) -> bool {
self.buffer.lock().is_empty()
}
pub fn len(&self) -> usize {
self.buffer.lock().len()
}
pub fn remove(&self, id: QueuedMessageId) -> Option<AgentMessage> {
let mut buf = self.buffer.lock();
let index = buf.iter().position(|item| item.id == id)?;
buf.remove(index).map(|item| item.message)
}
pub fn clear(&self) {
self.buffer.lock().clear();
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::agent::AgentMessage;
use crate::types::UserContent;
use std::sync::Arc;
fn make_msg(text: &str) -> AgentMessage {
AgentMessage::from(text)
}
#[test]
fn test_push_and_drain_all() {
let q = MessageQueue::new(QueueKind::Steering);
q.push(make_msg("a"));
q.push(make_msg("b"));
q.push(make_msg("c"));
let drained = q.drain_local(QueueMode::All);
assert_eq!(drained.len(), 3);
assert!(q.is_empty());
}
#[test]
fn test_push_and_drain_one_at_a_time() {
let q = MessageQueue::new(QueueKind::FollowUp);
q.push(make_msg("a"));
q.push(make_msg("b"));
let drained = q.drain_local(QueueMode::OneAtATime);
assert_eq!(drained.len(), 1);
assert_eq!(q.len(), 1);
let drained2 = q.drain_local(QueueMode::OneAtATime);
assert_eq!(drained2.len(), 1);
assert!(q.is_empty());
}
#[test]
fn test_drain_local_empty() {
let q = MessageQueue::new(QueueKind::Steering);
let drained = q.drain_local(QueueMode::All);
assert!(drained.is_empty());
let drained2 = q.drain_local(QueueMode::OneAtATime);
assert!(drained2.is_empty());
}
#[test]
fn test_clear() {
let q = MessageQueue::new(QueueKind::Steering);
q.push(make_msg("a"));
q.push(make_msg("b"));
assert_eq!(q.len(), 2);
q.clear();
assert!(q.is_empty());
assert_eq!(q.len(), 0);
}
#[test]
fn test_push_many() {
let q = MessageQueue::new(QueueKind::FollowUp);
q.push_many(vec![make_msg("a"), make_msg("b"), make_msg("c")]);
assert_eq!(q.len(), 3);
}
#[test]
fn test_push_many_bypasses_backpressure() {
let q = MessageQueue::new(QueueKind::FollowUp);
q.set_backpressure(BackpressureConfig {
max_depth: 1,
overflow: OverflowBehavior::Reject,
});
q.push_many(vec![make_msg("a"), make_msg("b"), make_msg("c")]);
assert_eq!(q.len(), 3);
let err = q.try_push(make_msg("d")).unwrap_err();
assert_eq!(err.current_depth, 3);
assert_eq!(err.max_depth, 1);
}
#[test]
fn test_remove_by_id_before_drain() {
let q = MessageQueue::new(QueueKind::Steering);
let id_a = q.push(make_msg("a"));
let id_b = q.push(make_msg("b"));
let id_c = q.push(make_msg("c"));
let removed = q.remove(id_b);
assert_eq!(removed, Some(make_msg("b")));
assert_eq!(q.len(), 2);
let drained = q.drain_local(QueueMode::All);
assert_eq!(drained, vec![make_msg("a"), make_msg("c")]);
assert!(q.remove(id_a).is_none());
assert!(q.remove(id_c).is_none());
}
#[test]
fn test_remove_unknown_id_returns_none() {
let q = MessageQueue::new(QueueKind::FollowUp);
let id = q.push(make_msg("a"));
assert_eq!(q.remove(QueuedMessageId(id.as_u64() + 1)), None);
assert_eq!(q.len(), 1);
}
#[test]
fn test_remove_after_drain_returns_none() {
let q = MessageQueue::new(QueueKind::Steering);
let id = q.push(make_msg("a"));
assert_eq!(q.drain_local(QueueMode::All), vec![make_msg("a")]);
assert!(q.remove(id).is_none());
}
#[test]
fn test_reject_then_remove_then_accept() {
let q = MessageQueue::new(QueueKind::FollowUp);
q.set_backpressure(BackpressureConfig {
max_depth: 1,
overflow: OverflowBehavior::Reject,
});
let id = q.try_push(make_msg("a")).unwrap();
assert!(q.try_push(make_msg("b")).is_err());
assert_eq!(q.remove(id), Some(make_msg("a")));
assert!(q.try_push(make_msg("b")).is_ok());
assert_eq!(q.drain_local(QueueMode::All), vec![make_msg("b")]);
}
#[test]
fn test_drop_oldest_makes_old_handle_uncancellable() {
let q = MessageQueue::new(QueueKind::Steering);
q.set_backpressure(BackpressureConfig {
max_depth: 1,
overflow: OverflowBehavior::DropOldest,
});
let old_id = q.try_push(make_msg("old")).unwrap();
let new_id = q.try_push(make_msg("new")).unwrap();
assert!(q.remove(old_id).is_none());
assert_eq!(q.remove(new_id), Some(make_msg("new")));
assert!(q.is_empty());
}
#[tokio::test]
async fn test_drain_with_supplier_all_mode() {
let q = MessageQueue::new(QueueKind::Steering);
q.push(make_msg("local"));
let supplier: GetQueuedMessagesFn =
Arc::new(|_signal| Box::pin(async { vec![AgentMessage::from("dynamic")] }));
let abort = AbortSignal::new();
let result = q.drain(QueueMode::All, &Some(supplier), abort).await;
assert_eq!(result.len(), 2); }
#[tokio::test]
async fn test_drain_one_at_a_time_local_first() {
let q = MessageQueue::new(QueueKind::Steering);
q.push(make_msg("local"));
let supplier: GetQueuedMessagesFn =
Arc::new(|_signal| Box::pin(async { vec![AgentMessage::from("dynamic")] }));
let abort = AbortSignal::new();
let result = q.drain(QueueMode::OneAtATime, &Some(supplier), abort).await;
assert_eq!(result.len(), 1);
}
#[tokio::test]
async fn test_drain_one_at_a_time_falls_to_supplier() {
let q = MessageQueue::new(QueueKind::Steering);
let supplier: GetQueuedMessagesFn = Arc::new(|_signal| {
Box::pin(async { vec![AgentMessage::from("d1"), AgentMessage::from("d2")] })
});
let abort = AbortSignal::new();
let result = q.drain(QueueMode::OneAtATime, &Some(supplier), abort).await;
assert_eq!(result.len(), 1);
}
#[tokio::test]
async fn test_drain_all_with_empty_local_uses_supplier_only() {
let q = MessageQueue::new(QueueKind::Steering);
let supplier: GetQueuedMessagesFn =
Arc::new(|_signal| Box::pin(async { vec![AgentMessage::from("dynamic")] }));
let abort = AbortSignal::new();
let result = q.drain(QueueMode::All, &Some(supplier), abort).await;
assert_eq!(result, vec![AgentMessage::from("dynamic")]);
assert!(q.is_empty());
}
#[tokio::test]
async fn test_drain_one_at_a_time_empty_supplier_returns_empty() {
let q = MessageQueue::new(QueueKind::Steering);
let supplier: GetQueuedMessagesFn = Arc::new(|_signal| Box::pin(async { Vec::new() }));
let abort = AbortSignal::new();
let result = q.drain(QueueMode::OneAtATime, &Some(supplier), abort).await;
assert!(result.is_empty());
assert!(q.is_empty());
}
#[tokio::test]
async fn test_drain_no_supplier() {
let q = MessageQueue::new(QueueKind::FollowUp);
q.push(make_msg("a"));
let abort = AbortSignal::new();
let result = q.drain(QueueMode::All, &None, abort).await;
assert_eq!(result.len(), 1);
}
#[test]
fn test_drain_strategy_all() {
let mut buf = VecDeque::from(vec![make_msg("a"), make_msg("b"), make_msg("c")]);
let strategy = DrainAll;
let result = strategy.select(&mut buf);
assert_eq!(result.len(), 3);
assert!(buf.is_empty());
}
#[test]
fn test_drain_strategy_one() {
let mut buf = VecDeque::from(vec![make_msg("a"), make_msg("b")]);
let strategy = DrainOne;
let result = strategy.select(&mut buf);
assert_eq!(result.len(), 1);
assert_eq!(buf.len(), 1);
}
#[test]
fn test_drain_strategy_batch() {
let mut buf = VecDeque::from(vec![
make_msg("a"),
make_msg("b"),
make_msg("c"),
make_msg("d"),
]);
let strategy = DrainBatch { max: 2 };
let result = strategy.select(&mut buf);
assert_eq!(result.len(), 2);
assert_eq!(buf.len(), 2);
}
#[test]
fn test_drain_with_strategy() {
let q = MessageQueue::new(QueueKind::Steering);
q.push_many(vec![make_msg("a"), make_msg("b"), make_msg("c")]);
let strategy = DrainBatch { max: 2 };
let result = q.drain_with_strategy(&strategy);
assert_eq!(result.len(), 2);
assert_eq!(q.len(), 1);
}
#[test]
fn test_backpressure_unlimited() {
let q = MessageQueue::new(QueueKind::Steering);
for i in 0..1000 {
assert!(q.try_push(make_msg(&format!("msg{}", i))).is_ok());
}
assert_eq!(q.len(), 1000);
}
#[test]
fn test_backpressure_reject() {
let q = MessageQueue::new(QueueKind::Steering);
q.set_backpressure(BackpressureConfig {
max_depth: 3,
overflow: OverflowBehavior::Reject,
});
assert!(q.try_push(make_msg("a")).is_ok());
assert!(q.try_push(make_msg("b")).is_ok());
assert!(q.try_push(make_msg("c")).is_ok());
let err = q.try_push(make_msg("d")).unwrap_err();
assert_eq!(err.current_depth, 3);
assert_eq!(err.max_depth, 3);
assert_eq!(q.len(), 3);
}
#[test]
fn test_backpressure_drop_oldest() {
let q = MessageQueue::new(QueueKind::Steering);
q.set_backpressure(BackpressureConfig {
max_depth: 3,
overflow: OverflowBehavior::DropOldest,
});
assert!(q.try_push(make_msg("a")).is_ok());
assert!(q.try_push(make_msg("b")).is_ok());
assert!(q.try_push(make_msg("c")).is_ok());
assert!(q.try_push(make_msg("d")).is_ok());
assert_eq!(q.len(), 3);
let msgs = q.drain_local(QueueMode::All);
assert_eq!(msgs.len(), 3);
let texts: Vec<&str> = msgs
.iter()
.filter_map(|m| match m {
AgentMessage::User(u) => match &u.content {
UserContent::Text(t) => Some(t.as_str()),
_ => None,
},
_ => None,
})
.collect();
assert_eq!(texts, vec!["b", "c", "d"]);
}
#[test]
fn test_queue_mode_to_strategy() {
let mut buf = VecDeque::from(vec![make_msg("a"), make_msg("b")]);
let strategy = QueueMode::All.to_strategy();
let result = strategy.select(&mut buf);
assert_eq!(result.len(), 2);
let mut buf2 = VecDeque::from(vec![make_msg("x"), make_msg("y")]);
let strategy2 = QueueMode::OneAtATime.to_strategy();
let result2 = strategy2.select(&mut buf2);
assert_eq!(result2.len(), 1);
assert_eq!(buf2.len(), 1);
}
#[test]
fn test_backpressure_zero_depth_with_reject_still_unlimited() {
let q = MessageQueue::new(QueueKind::Steering);
q.set_backpressure(BackpressureConfig {
max_depth: 0,
overflow: OverflowBehavior::Reject,
});
for i in 0..100 {
assert!(
q.try_push(make_msg(&format!("msg{}", i))).is_ok(),
"try_push should succeed when max_depth=0 even with Reject"
);
}
assert_eq!(q.len(), 100);
}
#[test]
fn test_backpressure_zero_depth_with_drop_oldest_still_unlimited() {
let q = MessageQueue::new(QueueKind::Steering);
q.set_backpressure(BackpressureConfig {
max_depth: 0,
overflow: OverflowBehavior::DropOldest,
});
for i in 0..100 {
assert!(
q.try_push(make_msg(&format!("msg{}", i))).is_ok(),
"try_push should succeed when max_depth=0 even with DropOldest"
);
}
assert_eq!(q.len(), 100);
let msgs = q.drain_local(QueueMode::All);
let texts: Vec<&str> = msgs
.iter()
.filter_map(|m| match m {
AgentMessage::User(u) => match &u.content {
UserContent::Text(t) => Some(t.as_str()),
_ => None,
},
_ => None,
})
.collect();
assert_eq!(texts.first(), Some(&"msg0"));
}
}