#![allow(clippy::all)]
use log::{debug, warn};
use std::collections::{BTreeMap, VecDeque};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::{Duration, Instant};
use tokio::sync::oneshot;
use super::priority::{Priority, RequestSource};
use crate::domain::EmbedRequest;
use crate::error::VecboostError;
#[derive(Debug)]
pub struct QueuedRequest {
pub request_id: String,
pub embed_request: EmbedRequest,
pub priority: Priority,
pub submitted_at: Instant,
pub timeout: Duration,
pub source: RequestSource,
pub response_tx: oneshot::Sender<Result<crate::domain::EmbedResponse, VecboostError>>,
}
pub struct PriorityRequestQueue {
queues: Arc<tokio::sync::RwLock<BTreeMap<Priority, VecDeque<QueuedRequest>>>>,
max_queue_size: usize,
current_size: Arc<AtomicUsize>,
}
impl PriorityRequestQueue {
pub fn new(max_queue_size: usize) -> Self {
debug!(
"Creating PriorityRequestQueue with max_size={}",
max_queue_size
);
Self {
queues: Arc::new(tokio::sync::RwLock::new(BTreeMap::new())),
max_queue_size,
current_size: Arc::new(AtomicUsize::new(0)),
}
}
pub async fn enqueue(&self, request: QueuedRequest) -> Result<(), VecboostError> {
loop {
let current_size = self.current_size.load(Ordering::Acquire);
if current_size >= self.max_queue_size {
return Err(VecboostError::RateLimitExceeded(
"Queue is full, request rejected".to_string(),
));
}
match self.current_size.compare_exchange_weak(
current_size,
current_size + 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
break;
}
Err(_) => {
continue;
}
}
}
let mut queues = self.queues.write().await;
let priority = request.priority;
let queue = queues.entry(priority).or_insert_with(VecDeque::new);
queue.push_back(request);
debug!(
"Request enqueued, priority={:?}, queue_size={}",
priority,
self.current_size.load(Ordering::Relaxed)
);
Ok(())
}
pub async fn dequeue(&self) -> Option<QueuedRequest> {
let mut queues = self.queues.write().await;
let now = Instant::now();
const AGING_THRESHOLD: Duration = Duration::from_secs(30);
for priority in [
Priority::Critical,
Priority::High,
Priority::Normal,
Priority::Low,
] {
if let Some(queue) = queues.get_mut(&priority) {
if priority != Priority::Low
&& let Some(front) = queue.front()
&& now.duration_since(front.submitted_at) > AGING_THRESHOLD
{
continue;
}
if let Some(request) = queue.pop_front() {
let new_size = self.current_size.fetch_sub(1, Ordering::Relaxed) - 1;
debug!(
"Request dequeued, priority={:?}, queue_size={}",
priority, new_size
);
if queue.is_empty() {
queues.remove(&priority);
}
return Some(request);
}
}
}
None
}
pub async fn peek_highest_priority(&self) -> Option<Priority> {
let queues = self.queues.read().await;
for priority in [
Priority::Critical,
Priority::High,
Priority::Normal,
Priority::Low,
] {
if let Some(queue) = queues.get(&priority) {
if !queue.is_empty() {
return Some(priority);
}
}
}
None
}
pub fn size(&self) -> usize {
self.current_size.load(Ordering::Relaxed)
}
pub async fn clear(&self) {
let mut queues = self.queues.write().await;
let cleared_count = queues.values().map(|q| q.len()).sum::<usize>();
queues.clear();
self.current_size.store(0, Ordering::Relaxed);
warn!("Queue cleared, {} requests discarded", cleared_count);
}
pub async fn size_by_priority(&self) -> Vec<(Priority, usize)> {
let queues = self.queues.read().await;
queues
.iter()
.map(|(priority, queue)| (*priority, queue.len()))
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_queue_creation() {
let queue = PriorityRequestQueue::new(100);
assert_eq!(queue.size(), 0);
}
#[tokio::test]
async fn test_enqueue_dequeue() {
let queue = PriorityRequestQueue::new(100);
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: "test-1".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
assert_eq!(queue.size(), 1);
let dequeued = queue.dequeue().await;
assert!(dequeued.is_some());
assert_eq!(queue.size(), 0);
}
#[tokio::test]
async fn test_priority_ordering() {
let queue = PriorityRequestQueue::new(100);
for (i, priority) in [
Priority::Low,
Priority::Critical,
Priority::Normal,
Priority::High,
]
.iter()
.enumerate()
{
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: format!("test-{}", i),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: *priority,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
}
assert_eq!(queue.dequeue().await.unwrap().priority, Priority::Critical);
assert_eq!(queue.dequeue().await.unwrap().priority, Priority::High);
assert_eq!(queue.dequeue().await.unwrap().priority, Priority::Normal);
assert_eq!(queue.dequeue().await.unwrap().priority, Priority::Low);
}
#[tokio::test]
async fn test_queue_full() {
let queue = PriorityRequestQueue::new(2);
let (tx1, _rx1) = oneshot::channel();
let request1 = QueuedRequest {
request_id: "test-1".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx1,
};
let (tx2, _rx2) = oneshot::channel();
let request2 = QueuedRequest {
request_id: "test-2".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx2,
};
queue.enqueue(request1).await.unwrap();
queue.enqueue(request2).await.unwrap();
let (tx3, _rx3) = oneshot::channel();
let request3 = QueuedRequest {
request_id: "test-3".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx3,
};
let result = queue.enqueue(request3).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_clear() {
let queue = PriorityRequestQueue::new(100);
for i in 0..10 {
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: format!("test-{}", i),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
}
assert_eq!(queue.size(), 10);
queue.clear().await;
assert_eq!(queue.size(), 0);
}
#[tokio::test]
async fn test_aging_prevents_low_priority_starvation() {
let queue = PriorityRequestQueue::new(100);
let (tx1, _rx1) = oneshot::channel();
let critical_req = QueuedRequest {
request_id: "critical-1".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Critical,
submitted_at: Instant::now(),
timeout: Duration::from_secs(60),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx1,
};
queue.enqueue(critical_req).await.unwrap();
let (tx2, _rx2) = oneshot::channel();
let low_req = QueuedRequest {
request_id: "low-1".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Low,
submitted_at: Instant::now(),
timeout: Duration::from_secs(60),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx2,
};
queue.enqueue(low_req).await.unwrap();
tokio::time::sleep(Duration::from_secs(31)).await;
let dequeued = queue.dequeue().await.unwrap();
assert_eq!(
dequeued.priority,
Priority::Low,
"aged Critical should be skipped, Low should be dequeued first"
);
}
#[tokio::test]
async fn test_peek_highest_priority_empty_queue_returns_none() {
let queue = PriorityRequestQueue::new(100);
let result = queue.peek_highest_priority().await;
assert!(result.is_none());
}
#[tokio::test]
async fn test_peek_highest_priority_returns_critical() {
let queue = PriorityRequestQueue::new(100);
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: "test-1".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Critical,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
let result = queue.peek_highest_priority().await;
assert_eq!(result, Some(Priority::Critical));
}
#[tokio::test]
async fn test_peek_highest_priority_returns_low() {
let queue = PriorityRequestQueue::new(100);
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: "test-1".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Low,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
let result = queue.peek_highest_priority().await;
assert_eq!(result, Some(Priority::Low));
}
#[tokio::test]
async fn test_peek_highest_priority_after_dequeue() {
let queue = PriorityRequestQueue::new(100);
for priority in [Priority::High, Priority::Low] {
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: format!("test-{:?}", priority),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
}
assert_eq!(queue.peek_highest_priority().await, Some(Priority::High));
queue.dequeue().await.unwrap();
assert_eq!(queue.peek_highest_priority().await, Some(Priority::Low));
queue.dequeue().await.unwrap();
assert!(queue.peek_highest_priority().await.is_none());
}
#[tokio::test]
async fn test_size_by_priority_empty_queue() {
let queue = PriorityRequestQueue::new(100);
let result = queue.size_by_priority().await;
assert!(result.is_empty());
}
#[tokio::test]
async fn test_size_by_priority_single_priority() {
let queue = PriorityRequestQueue::new(100);
for i in 0..3 {
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: format!("test-{}", i),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
}
let result = queue.size_by_priority().await;
assert_eq!(result.len(), 1);
assert_eq!(result[0], (Priority::Normal, 3));
}
#[tokio::test]
async fn test_size_by_priority_multiple_priorities() {
let queue = PriorityRequestQueue::new(100);
let priorities_with_counts = [
(Priority::Critical, 2),
(Priority::High, 1),
(Priority::Low, 3),
];
for (priority, count) in priorities_with_counts {
for i in 0..count {
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: format!("test-{:?}-{}", priority, i),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
}
}
let result = queue.size_by_priority().await;
let total: usize = result.iter().map(|(_, c)| c).sum();
assert_eq!(total, 6);
}
#[tokio::test]
async fn test_dequeue_empty_queue_returns_none() {
let queue = PriorityRequestQueue::new(100);
let result = queue.dequeue().await;
assert!(result.is_none());
}
#[tokio::test]
async fn test_dequeue_all_then_empty() {
let queue = PriorityRequestQueue::new(100);
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: "test-1".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
assert!(queue.dequeue().await.is_some());
assert!(queue.dequeue().await.is_none());
assert_eq!(queue.size(), 0);
}
#[tokio::test]
async fn test_clear_empty_queue() {
let queue = PriorityRequestQueue::new(100);
queue.clear().await;
assert_eq!(queue.size(), 0);
}
#[tokio::test]
async fn test_size_reflects_enqueue_and_dequeue() {
let queue = PriorityRequestQueue::new(100);
assert_eq!(queue.size(), 0);
for i in 0..5 {
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: format!("test-{}", i),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
}
assert_eq!(queue.size(), 5);
for _ in 0..3 {
queue.dequeue().await.unwrap();
}
assert_eq!(queue.size(), 2);
}
#[tokio::test]
async fn test_enqueue_max_size_zero_always_rejects() {
let queue = PriorityRequestQueue::new(0);
let (tx, _rx) = oneshot::channel();
let request = QueuedRequest {
request_id: "test-1".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
let result = queue.enqueue(request).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_aging_does_not_skip_low_priority() {
let queue = PriorityRequestQueue::new(100);
let (tx, _rx) = oneshot::channel();
let low_req = QueuedRequest {
request_id: "low-1".to_string(),
embed_request: EmbedRequest {
text: "test".to_string(),
normalize: Some(true),
},
priority: Priority::Low,
submitted_at: Instant::now(),
timeout: Duration::from_secs(60),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
response_tx: tx,
};
queue.enqueue(low_req).await.unwrap();
let dequeued = queue.dequeue().await.unwrap();
assert_eq!(dequeued.priority, Priority::Low);
}
}