mod attempt;
pub mod auth;
pub mod http;
pub mod idempotency;
pub mod metrics;
pub mod quota;
#[cfg(feature = "redis-coordination")]
pub mod distributed;
use std::collections::{BTreeMap, HashMap, VecDeque};
use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use async_trait::async_trait;
use serde_json::Value;
use tokio::sync::{mpsc, oneshot, Notify, Semaphore};
use tokio::task::JoinHandle;
use tokio::time::Instant;
use crate::proxy::ratelimit::{RateKey, RateLimiter, RetryAfter};
pub type Tier = u8;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct PriorityKey {
tier: Tier,
seqno: u64,
}
pub struct GatewayRequest {
pub provider: String,
pub tier: Tier,
pub permits: u32,
pub payload: Value,
}
#[derive(Debug)]
pub enum GatewayError {
Overloaded(Duration),
Timeout,
Upstream(String),
Shutdown,
}
impl std::fmt::Display for GatewayError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
GatewayError::Overloaded(d) => write!(f, "gateway overloaded, retry in {d:?}"),
GatewayError::Timeout => write!(f, "gateway queue wait timed out"),
GatewayError::Upstream(m) => write!(f, "upstream error: {m}"),
GatewayError::Shutdown => write!(f, "gateway shutting down"),
}
}
}
impl std::error::Error for GatewayError {}
pub struct DispatchError {
pub message: String,
pub retry_after: Option<Duration>,
}
impl DispatchError {
pub fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
retry_after: None,
}
}
}
#[async_trait]
pub trait Dispatch: Send + Sync {
async fn dispatch(&self, provider: &str, payload: Value) -> Result<Value, DispatchError>;
async fn dispatch_with_policy(
&self,
provider: &str,
payload: Value,
_policy_context: crate::policy::DispatchPolicyContext,
) -> Result<Value, DispatchError> {
self.dispatch(provider, payload).await
}
async fn dispatch_stream(
&self,
_provider: &str,
_payload: Value,
) -> Result<ChunkStream, DispatchError> {
Err(DispatchError::new(
"streaming not supported by this dispatcher",
))
}
async fn dispatch_stream_with_policy(
&self,
provider: &str,
payload: Value,
_policy_context: crate::policy::DispatchPolicyContext,
) -> Result<ChunkStream, DispatchError> {
self.dispatch_stream(provider, payload).await
}
}
pub type StreamChunk = Result<String, GatewayError>;
pub type ChunkStream = std::pin::Pin<Box<dyn futures::Stream<Item = StreamChunk> + Send>>;
enum Delivery {
Unary(oneshot::Sender<Result<Value, GatewayError>>),
Stream(oneshot::Sender<Result<mpsc::Receiver<StreamChunk>, GatewayError>>),
}
pub struct Job {
key: PriorityKey,
permits: u32,
payload: Value,
policy_context: Option<crate::policy::DispatchPolicyContext>,
delivery: Delivery,
started: oneshot::Sender<()>,
enqueued_at: Instant,
}
impl Job {
fn is_cancelled(&self) -> bool {
match &self.delivery {
Delivery::Unary(tx) => tx.is_closed(),
Delivery::Stream(tx) => tx.is_closed(),
}
}
}
#[async_trait]
pub trait RequestQueue: Send + Sync {
fn enqueue(&self, job: Job) -> Result<(), Job>;
async fn dequeue(&self) -> Job;
fn requeue(&self, job: Job);
async fn notified(&self);
fn depth(&self) -> usize;
}
pub struct InMemoryQueue {
inner: Mutex<QueueInner>,
notify: Notify,
max_depth: usize,
aging_step: Duration,
max_boost: u32,
}
#[derive(Default)]
struct QueueInner {
tiers: BTreeMap<Tier, VecDeque<Job>>,
len: usize,
}
impl InMemoryQueue {
pub fn new(max_depth: usize) -> Self {
Self::with_aging(max_depth, Duration::from_secs(5), 16)
}
pub fn with_aging(max_depth: usize, aging_step: Duration, max_boost: u32) -> Self {
Self {
inner: Mutex::new(QueueInner::default()),
notify: Notify::new(),
max_depth: max_depth.max(1),
aging_step: aging_step.max(Duration::from_millis(1)),
max_boost,
}
}
fn effective_priority(&self, tier: Tier, enqueued_at: Instant, now: Instant) -> i64 {
let waited = now.saturating_duration_since(enqueued_at).as_secs_f64();
let boost = (waited / self.aging_step.as_secs_f64()).floor() as i64;
tier as i64 + boost.clamp(0, self.max_boost as i64)
}
fn pop_best(&self, inner: &mut QueueInner, now: Instant) -> Option<Job> {
let tiers: Vec<Tier> = inner.tiers.keys().copied().collect();
let mut best: Option<(i64, u64, Tier)> = None; for tier in tiers {
let dq = inner.tiers.get_mut(&tier).unwrap();
while dq.front().map(|j| j.is_cancelled()).unwrap_or(false) {
dq.pop_front();
inner.len -= 1;
}
if let Some(front) = dq.front() {
let eff = self.effective_priority(tier, front.enqueued_at, now);
let seqno = front.key.seqno;
let better = match best {
None => true,
Some((be, bs, _)) => eff > be || (eff == be && seqno < bs),
};
if better {
best = Some((eff, seqno, tier));
}
}
}
inner.tiers.retain(|_, dq| !dq.is_empty());
let (_, _, tier) = best?;
let dq = inner.tiers.get_mut(&tier).unwrap();
let job = dq.pop_front();
if dq.is_empty() {
inner.tiers.remove(&tier);
}
if job.is_some() {
inner.len -= 1;
}
job
}
}
#[async_trait]
impl RequestQueue for InMemoryQueue {
fn enqueue(&self, job: Job) -> Result<(), Job> {
{
let mut inner = self.inner.lock().unwrap();
if inner.len >= self.max_depth {
return Err(job);
}
inner.len += 1;
inner.tiers.entry(job.key.tier).or_default().push_back(job);
}
self.notify.notify_one();
Ok(())
}
async fn dequeue(&self) -> Job {
loop {
{
let mut inner = self.inner.lock().unwrap();
let now = Instant::now();
if let Some(job) = self.pop_best(&mut inner, now) {
return job;
}
}
self.notify.notified().await;
}
}
fn requeue(&self, job: Job) {
let mut inner = self.inner.lock().unwrap();
inner.len += 1;
inner.tiers.entry(job.key.tier).or_default().push_front(job);
}
async fn notified(&self) {
self.notify.notified().await;
}
fn depth(&self) -> usize {
self.inner.lock().unwrap().len
}
}
#[derive(Clone)]
pub struct GatewayConfig {
pub max_queue_depth: usize,
pub max_wait: Duration,
pub overloaded_retry_after: Duration,
pub max_concurrency_per_provider: usize,
pub aging_step: Duration,
pub max_boost: u32,
pub request_timeout: Duration,
pub lease_timeout: Duration,
pub max_attempts: u32,
}
impl Default for GatewayConfig {
fn default() -> Self {
Self {
max_queue_depth: 10_000,
max_wait: Duration::from_secs(30),
overloaded_retry_after: Duration::from_secs(1),
max_concurrency_per_provider: 256,
aging_step: Duration::from_secs(5),
max_boost: 16,
request_timeout: Duration::from_secs(120),
lease_timeout: Duration::from_secs(60),
max_attempts: 5,
}
}
}
impl GatewayConfig {
pub fn from_env() -> Self {
let d = Self::default();
let usize_env = |k: &str, fallback: usize| {
std::env::var(k)
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(fallback)
};
Self {
max_queue_depth: usize_env("LLMSHIM_GATEWAY_QUEUE_DEPTH", d.max_queue_depth),
max_wait: std::env::var("LLMSHIM_GATEWAY_MAX_WAIT_MS")
.ok()
.and_then(|v| v.parse().ok())
.map(Duration::from_millis)
.unwrap_or(d.max_wait),
overloaded_retry_after: d.overloaded_retry_after,
max_concurrency_per_provider: usize_env(
"LLMSHIM_GATEWAY_MAX_CONCURRENCY",
d.max_concurrency_per_provider,
),
aging_step: std::env::var("LLMSHIM_GATEWAY_AGING_STEP_MS")
.ok()
.and_then(|v| v.parse().ok())
.map(Duration::from_millis)
.unwrap_or(d.aging_step),
max_boost: std::env::var("LLMSHIM_GATEWAY_MAX_BOOST")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(d.max_boost),
request_timeout: std::env::var("LLMSHIM_GATEWAY_REQUEST_TIMEOUT_MS")
.ok()
.and_then(|v| v.parse().ok())
.map(Duration::from_millis)
.unwrap_or(d.request_timeout),
lease_timeout: std::env::var("LLMSHIM_GATEWAY_LEASE_TIMEOUT_MS")
.ok()
.and_then(|v| v.parse().ok())
.map(Duration::from_millis)
.unwrap_or(d.lease_timeout),
max_attempts: std::env::var("LLMSHIM_GATEWAY_MAX_ATTEMPTS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(d.max_attempts),
}
}
}
struct Lane {
queue: Arc<dyn RequestQueue>,
}
pub struct Scheduler {
limiter: Arc<dyn RateLimiter>,
dispatch: Arc<dyn Dispatch>,
config: GatewayConfig,
lanes: Mutex<HashMap<String, Lane>>,
seq: AtomicU64,
handles: Mutex<Vec<JoinHandle<()>>>,
}
impl Scheduler {
pub fn new(
config: GatewayConfig,
limiter: Arc<dyn RateLimiter>,
dispatch: Arc<dyn Dispatch>,
) -> Arc<Self> {
Arc::new(Self {
limiter,
dispatch,
config,
lanes: Mutex::new(HashMap::new()),
seq: AtomicU64::new(0),
handles: Mutex::new(Vec::new()),
})
}
pub async fn submit(self: &Arc<Self>, req: GatewayRequest) -> Result<Value, GatewayError> {
self.submit_inner(req, None).await
}
async fn submit_inner(
self: &Arc<Self>,
req: GatewayRequest,
policy_context: Option<crate::policy::DispatchPolicyContext>,
) -> Result<Value, GatewayError> {
let provider = req.provider.clone();
let tier = req.tier;
let queue = self.lane_for(&provider);
let seqno = self.seq.fetch_add(1, AtomicOrdering::Relaxed);
let (tx, rx) = oneshot::channel();
let (started_tx, started_rx) = oneshot::channel();
let job = Job {
key: PriorityKey {
tier: req.tier,
seqno,
},
permits: req.permits.max(1),
payload: req.payload,
policy_context,
delivery: Delivery::Unary(tx),
started: started_tx,
enqueued_at: Instant::now(),
};
if queue.enqueue(job).is_err() {
metrics::incr(
metrics::REJECTED,
&[("provider", &provider), ("reason", "overloaded")],
);
return Err(GatewayError::Overloaded(self.config.overloaded_retry_after));
}
metrics::incr(
metrics::REQUESTS,
&[
("provider", &provider),
("tier", &tier.to_string()),
("mode", "unary"),
],
);
let result = tokio::select! {
biased;
started = started_rx => match started {
Ok(()) | Err(_) => match rx.await {
Ok(result) => result,
Err(_) => Err(GatewayError::Shutdown),
},
},
_ = tokio::time::sleep(self.config.max_wait) => Err(GatewayError::Timeout),
};
if matches!(result, Err(GatewayError::Timeout)) {
metrics::incr(
metrics::REJECTED,
&[("provider", &provider), ("reason", "timeout")],
);
}
result
}
pub(crate) async fn submit_with_policy(
self: &Arc<Self>,
req: GatewayRequest,
policy_context: crate::policy::DispatchPolicyContext,
) -> Result<Value, GatewayError> {
self.submit_inner(req, Some(policy_context)).await
}
pub async fn submit_stream(
self: &Arc<Self>,
req: GatewayRequest,
) -> Result<mpsc::Receiver<StreamChunk>, GatewayError> {
self.submit_stream_inner(req, None).await
}
pub(crate) async fn submit_stream_with_policy(
self: &Arc<Self>,
req: GatewayRequest,
policy_context: crate::policy::DispatchPolicyContext,
) -> Result<mpsc::Receiver<StreamChunk>, GatewayError> {
self.submit_stream_inner(req, Some(policy_context)).await
}
async fn submit_stream_inner(
self: &Arc<Self>,
req: GatewayRequest,
policy_context: Option<crate::policy::DispatchPolicyContext>,
) -> Result<mpsc::Receiver<StreamChunk>, GatewayError> {
let provider = req.provider.clone();
let tier = req.tier;
let queue = self.lane_for(&provider);
let seqno = self.seq.fetch_add(1, AtomicOrdering::Relaxed);
let (result_tx, result_rx) = oneshot::channel();
let (started_tx, started_rx) = oneshot::channel();
let job = Job {
key: PriorityKey {
tier: req.tier,
seqno,
},
permits: req.permits.max(1),
payload: req.payload,
policy_context,
delivery: Delivery::Stream(result_tx),
started: started_tx,
enqueued_at: Instant::now(),
};
if queue.enqueue(job).is_err() {
metrics::incr(
metrics::REJECTED,
&[("provider", &provider), ("reason", "overloaded")],
);
return Err(GatewayError::Overloaded(self.config.overloaded_retry_after));
}
metrics::incr(
metrics::REQUESTS,
&[
("provider", &provider),
("tier", &tier.to_string()),
("mode", "stream"),
],
);
let result = tokio::select! {
biased;
started = started_rx => match started {
Ok(()) | Err(_) => match result_rx.await {
Ok(result) => result,
Err(_) => Err(GatewayError::Shutdown),
},
},
_ = tokio::time::sleep(self.config.max_wait) => Err(GatewayError::Timeout),
};
if matches!(result, Err(GatewayError::Timeout)) {
metrics::incr(
metrics::REJECTED,
&[("provider", &provider), ("reason", "timeout")],
);
}
result
}
pub fn lane_depths(&self) -> Vec<(String, usize)> {
self.lanes
.lock()
.unwrap()
.iter()
.map(|(p, lane)| (p.clone(), lane.queue.depth()))
.collect()
}
pub fn queue_depth(&self, provider: &str) -> usize {
self.lanes
.lock()
.unwrap()
.get(provider)
.map(|l| l.queue.depth())
.unwrap_or(0)
}
fn lane_for(self: &Arc<Self>, provider: &str) -> Arc<dyn RequestQueue> {
let mut lanes = self.lanes.lock().unwrap();
if let Some(lane) = lanes.get(provider) {
return lane.queue.clone();
}
let queue: Arc<dyn RequestQueue> = Arc::new(InMemoryQueue::with_aging(
self.config.max_queue_depth,
self.config.aging_step,
self.config.max_boost,
));
lanes.insert(
provider.to_string(),
Lane {
queue: queue.clone(),
},
);
let handle = tokio::spawn(dispatcher_loop(
provider.to_string(),
queue.clone(),
self.limiter.clone(),
self.dispatch.clone(),
self.config.max_concurrency_per_provider,
));
self.handles.lock().unwrap().push(handle);
queue
}
}
impl Drop for Scheduler {
fn drop(&mut self) {
for handle in self.handles.lock().unwrap().drain(..) {
handle.abort();
}
}
}
async fn penalize_if_429(limiter: &Arc<dyn RateLimiter>, provider: &str, err: &DispatchError) {
if let Some(retry_after) = err.retry_after {
limiter
.penalize(&RateKey::provider(provider.to_string()), retry_after)
.await;
}
}
async fn dispatcher_loop(
provider: String,
queue: Arc<dyn RequestQueue>,
limiter: Arc<dyn RateLimiter>,
dispatch: Arc<dyn Dispatch>,
max_concurrency: usize,
) {
let key = RateKey::provider(provider.clone());
let sem = Arc::new(Semaphore::new(max_concurrency.max(1)));
loop {
let job = queue.dequeue().await;
if job.is_cancelled() {
continue; }
let policy_gated = job.policy_context.is_some();
let permit = if policy_gated {
None
} else {
match sem.clone().acquire_owned().await {
Ok(permit) => Some(permit),
Err(_) => break,
}
};
if job.is_cancelled() {
continue; }
let rate_admission = if policy_gated {
Ok(())
} else {
limiter.acquire(&key, job.permits).await
};
match rate_admission {
Ok(()) => {
if job.is_cancelled() {
continue; }
metrics::observe_ms(
metrics::QUEUE_WAIT,
&[("provider", &provider)],
job.enqueued_at.elapsed().as_millis() as f64,
);
let Job {
payload,
policy_context,
delivery,
started,
..
} = job;
let _ = started.send(());
let dispatch = dispatch.clone();
let limiter = limiter.clone();
let provider = provider.clone();
tokio::spawn(async move {
let _permit = permit;
let _inflight = metrics::inflight(&provider);
let started_at = Instant::now();
let plabels: &[(&str, &str)] = &[("provider", &provider)];
match delivery {
Delivery::Unary(tx) => match match policy_context {
Some(context) => {
dispatch
.dispatch_with_policy(&provider, payload, context)
.await
}
None => dispatch.dispatch(&provider, payload).await,
} {
Ok(value) => {
metrics::incr(metrics::DISPATCHED, plabels);
metrics::observe_ms(
metrics::UPSTREAM_LATENCY,
plabels,
started_at.elapsed().as_millis() as f64,
);
let _ = tx.send(Ok(value));
}
Err(err) => {
metrics::incr(
metrics::REJECTED,
&[("provider", &provider), ("reason", "upstream")],
);
if !policy_gated {
penalize_if_429(&limiter, &provider, &err).await;
}
let _ = tx.send(Err(GatewayError::Upstream(err.message)));
}
},
Delivery::Stream(tx) => {
let opened = match policy_context {
Some(context) => {
dispatch
.dispatch_stream_with_policy(&provider, payload, context)
.await
}
None => dispatch.dispatch_stream(&provider, payload).await,
};
match opened {
Ok(mut upstream) => {
metrics::incr(metrics::DISPATCHED, plabels);
let (chunk_tx, chunk_rx) = mpsc::channel(16);
if tx.send(Ok(chunk_rx)).is_err() {
return;
}
use futures::StreamExt;
while let Some(item) = upstream.next().await {
if chunk_tx.send(item).await.is_err() {
break;
}
}
metrics::observe_ms(
metrics::UPSTREAM_LATENCY,
plabels,
started_at.elapsed().as_millis() as f64,
);
}
Err(err) => {
metrics::incr(
metrics::REJECTED,
&[("provider", &provider), ("reason", "upstream")],
);
if !policy_gated {
penalize_if_429(&limiter, &provider, &err).await;
}
let _ = tx.send(Err(GatewayError::Upstream(err.message)));
}
}
}
}
});
}
Err(RetryAfter(wait)) => {
drop(permit);
queue.requeue(job);
tokio::select! {
_ = tokio::time::sleep(wait) => {}
_ = queue.notified() => {}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use std::collections::HashMap as Map;
struct FakeLimiter {
permits: Mutex<Map<String, i64>>,
default: i64,
}
impl FakeLimiter {
fn new(default: i64) -> Self {
Self {
permits: Mutex::new(Map::new()),
default,
}
}
fn set(&self, provider: &str, n: i64) {
self.permits.lock().unwrap().insert(provider.to_string(), n);
}
}
#[async_trait]
impl RateLimiter for FakeLimiter {
async fn acquire(&self, key: &RateKey, permits: u32) -> Result<(), RetryAfter> {
let mut map = self.permits.lock().unwrap();
let bal = map.entry(key.provider.clone()).or_insert(self.default);
if *bal >= permits as i64 {
*bal -= permits as i64;
Ok(())
} else {
Err(RetryAfter(Duration::from_secs(1)))
}
}
async fn penalize(&self, _key: &RateKey, _retry_after: Duration) {}
}
struct RecordingDispatch {
order: Arc<Mutex<Vec<u64>>>,
}
#[async_trait]
impl Dispatch for RecordingDispatch {
async fn dispatch(&self, _provider: &str, payload: Value) -> Result<Value, DispatchError> {
let id = payload["id"].as_u64().unwrap();
self.order.lock().unwrap().push(id);
Ok(json!({ "id": id }))
}
async fn dispatch_stream(
&self,
_provider: &str,
payload: Value,
) -> Result<ChunkStream, DispatchError> {
let id = payload["id"].as_u64().unwrap();
self.order.lock().unwrap().push(id);
let chunks: Vec<StreamChunk> = (0..3).map(|n| Ok(format!("{id}:{n}"))).collect();
Ok(Box::pin(futures::stream::iter(chunks)))
}
}
fn scheduler(
limiter: Arc<FakeLimiter>,
order: Arc<Mutex<Vec<u64>>>,
config: GatewayConfig,
) -> Arc<Scheduler> {
Scheduler::new(config, limiter, Arc::new(RecordingDispatch { order }))
}
async fn yield_many() {
for _ in 0..50 {
tokio::task::yield_now().await;
}
}
async fn wait_for_depth(sched: &Arc<Scheduler>, provider: &str, target: usize) {
for _ in 0..10_000 {
if sched.queue_depth(provider) >= target {
return;
}
tokio::task::yield_now().await;
}
panic!("queue for {provider} never reached depth {target} (dispatcher gated?)");
}
async fn enqueue_ordered(
sched: &Arc<Scheduler>,
provider: &str,
items: &[(u64, Tier)],
) -> Vec<JoinHandle<Result<Value, GatewayError>>> {
let mut handles = Vec::new();
for (i, &(id, tier)) in items.iter().enumerate() {
let s = sched.clone();
let p = provider.to_string();
handles.push(tokio::spawn(async move {
s.submit(GatewayRequest {
provider: p,
tier,
permits: 1,
payload: json!({ "id": id }),
})
.await
}));
wait_for_depth(sched, provider, i + 1).await;
}
handles
}
fn dummy_job(tier: Tier, seqno: u64) -> (Job, oneshot::Receiver<Result<Value, GatewayError>>) {
let (tx, rx) = oneshot::channel();
let (started, _started_rx) = oneshot::channel();
(
Job {
key: PriorityKey { tier, seqno },
permits: 1,
payload: json!({ "id": seqno }),
policy_context: None,
delivery: Delivery::Unary(tx),
started,
enqueued_at: Instant::now(),
},
rx,
)
}
#[tokio::test]
async fn queue_dequeues_highest_tier_then_fifo() {
let q = InMemoryQueue::new(100);
let mut keep = Vec::new();
for (tier, seq) in [(1u8, 0u64), (3, 1), (1, 2), (2, 3), (3, 4)] {
let (job, rx) = dummy_job(tier, seq);
keep.push(rx);
assert!(q.enqueue(job).is_ok());
}
let mut got = Vec::new();
for _ in 0..5 {
got.push(q.dequeue().await.key);
}
let order: Vec<(u8, u64)> = got.iter().map(|k| (k.tier, k.seqno)).collect();
assert_eq!(order, vec![(3, 1), (3, 4), (2, 3), (1, 0), (1, 2)]);
}
#[tokio::test]
async fn queue_sheds_when_full() {
let q = InMemoryQueue::new(2);
let (j1, _r1) = dummy_job(0, 0);
let (j2, _r2) = dummy_job(0, 1);
let (j3, _r3) = dummy_job(0, 2);
assert!(q.enqueue(j1).is_ok());
assert!(q.enqueue(j2).is_ok());
assert!(q.enqueue(j3).is_err(), "third enqueue should shed");
assert_eq!(q.depth(), 2);
}
#[tokio::test(start_paused = true)]
async fn max_wait_bounds_queue_time_not_upstream_call() {
struct SlowDispatch;
#[async_trait]
impl Dispatch for SlowDispatch {
async fn dispatch(
&self,
_provider: &str,
payload: Value,
) -> Result<Value, DispatchError> {
tokio::time::sleep(Duration::from_secs(60)).await;
Ok(payload)
}
}
let limiter = Arc::new(FakeLimiter::new(1000)); let config = GatewayConfig {
max_wait: Duration::from_secs(5),
..Default::default()
};
let sched = Scheduler::new(config, limiter, Arc::new(SlowDispatch));
let s = sched.clone();
let handle = tokio::spawn(async move {
s.submit(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({ "id": 7 }),
})
.await
});
yield_many().await;
tokio::time::advance(Duration::from_secs(61)).await;
yield_many().await;
let result = handle.await.unwrap();
assert!(
matches!(&result, Ok(v) if v["id"] == 7),
"slow upstream call must complete, not time out: {result:?}"
);
}
#[tokio::test(start_paused = true)]
async fn dispatches_in_priority_then_fifo_order() {
let limiter = Arc::new(FakeLimiter::new(0)); let order = Arc::new(Mutex::new(Vec::new()));
let sched = scheduler(limiter.clone(), order.clone(), GatewayConfig::default());
let handles = enqueue_ordered(&sched, "p", &[(1, 1), (2, 3), (3, 1), (4, 2)]).await;
limiter.set("p", 100);
tokio::time::advance(Duration::from_secs(1)).await;
yield_many().await;
for h in handles {
h.await.unwrap().unwrap();
}
assert_eq!(*order.lock().unwrap(), vec![2, 4, 1, 3]);
}
#[tokio::test(start_paused = true)]
async fn aging_rescues_starved_low_tier() {
let limiter = Arc::new(FakeLimiter::new(0)); let order = Arc::new(Mutex::new(Vec::new()));
let config = GatewayConfig {
max_wait: Duration::from_secs(300),
..Default::default()
};
let sched = scheduler(limiter.clone(), order.clone(), config);
let mut handles = enqueue_ordered(&sched, "p", &[(1, 0)]).await;
tokio::time::advance(Duration::from_secs(31)).await;
yield_many().await;
let hi = {
let s = sched.clone();
tokio::spawn(async move {
s.submit(GatewayRequest {
provider: "p".into(),
tier: 5,
permits: 1,
payload: json!({ "id": 2 }),
})
.await
})
};
wait_for_depth(&sched, "p", 2).await;
handles.push(hi);
limiter.set("p", 100);
tokio::time::advance(Duration::from_secs(1)).await;
yield_many().await;
for h in handles {
h.await.unwrap().unwrap();
}
assert_eq!(*order.lock().unwrap(), vec![1, 2]);
}
#[tokio::test(start_paused = true)]
async fn gates_on_rate_limit_then_dispatches_on_refill() {
let limiter = Arc::new(FakeLimiter::new(0));
let order = Arc::new(Mutex::new(Vec::new()));
let sched = scheduler(limiter.clone(), order.clone(), GatewayConfig::default());
let handles = enqueue_ordered(&sched, "p", &[(1, 0), (2, 0)]).await;
limiter.set("p", 1);
tokio::time::advance(Duration::from_secs(1)).await;
yield_many().await;
assert_eq!(*order.lock().unwrap(), vec![1]);
limiter.set("p", 1);
tokio::time::advance(Duration::from_secs(1)).await;
yield_many().await;
assert_eq!(*order.lock().unwrap(), vec![1, 2]);
for h in handles {
h.await.unwrap().unwrap();
}
}
#[tokio::test(start_paused = true)]
async fn per_provider_independence() {
let limiter = Arc::new(FakeLimiter::new(0));
let order = Arc::new(Mutex::new(Vec::new()));
let sched = scheduler(limiter.clone(), order.clone(), GatewayConfig::default());
limiter.set("slow", 0); limiter.set("fast", 100);
let slow = {
let s = sched.clone();
tokio::spawn(async move {
s.submit(GatewayRequest {
provider: "slow".into(),
tier: 0,
permits: 1,
payload: json!({ "id": 99 }),
})
.await
})
};
let fast = {
let s = sched.clone();
tokio::spawn(async move {
s.submit(GatewayRequest {
provider: "fast".into(),
tier: 0,
permits: 1,
payload: json!({ "id": 1 }),
})
.await
})
};
yield_many().await;
assert_eq!(*order.lock().unwrap(), vec![1]);
fast.await.unwrap().unwrap();
slow.abort();
}
#[tokio::test(start_paused = true)]
async fn times_out_when_never_dispatched() {
let limiter = Arc::new(FakeLimiter::new(0)); let order = Arc::new(Mutex::new(Vec::new()));
let config = GatewayConfig {
max_wait: Duration::from_secs(2),
..Default::default()
};
let sched = scheduler(limiter.clone(), order.clone(), config);
let s = sched.clone();
let handle = tokio::spawn(async move {
s.submit(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({ "id": 1 }),
})
.await
});
yield_many().await;
tokio::time::advance(Duration::from_secs(2)).await;
yield_many().await;
let result = handle.await.unwrap();
assert!(matches!(result, Err(GatewayError::Timeout)));
assert!(order.lock().unwrap().is_empty(), "nothing should dispatch");
}
#[tokio::test(start_paused = true)]
async fn cancelled_job_is_not_dispatched() {
let limiter = Arc::new(FakeLimiter::new(0));
let order = Arc::new(Mutex::new(Vec::new()));
let sched = scheduler(limiter.clone(), order.clone(), GatewayConfig::default());
let s = sched.clone();
let handle = tokio::spawn(async move {
s.submit(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({ "id": 1 }),
})
.await
});
wait_for_depth(&sched, "p", 1).await;
handle.abort();
yield_many().await;
limiter.set("p", 100);
tokio::time::advance(Duration::from_secs(1)).await;
yield_many().await;
assert!(
order.lock().unwrap().is_empty(),
"cancelled job must not dispatch"
);
}
#[tokio::test(start_paused = true)]
async fn sheds_when_queue_full() {
let limiter = Arc::new(FakeLimiter::new(0)); let order = Arc::new(Mutex::new(Vec::new()));
let config = GatewayConfig {
max_queue_depth: 2,
..Default::default()
};
let sched = scheduler(limiter.clone(), order.clone(), config);
let _held = enqueue_ordered(&sched, "p", &[(1, 0), (2, 0)]).await;
let result = sched
.submit(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({ "id": 3 }),
})
.await;
assert!(matches!(result, Err(GatewayError::Overloaded(_))));
}
#[tokio::test(start_paused = true)]
async fn submit_stream_delivers_chunks_in_order() {
let limiter = Arc::new(FakeLimiter::new(100)); let order = Arc::new(Mutex::new(Vec::new()));
let sched = scheduler(limiter, order, GatewayConfig::default());
let mut rx = sched
.submit_stream(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({ "id": 7 }),
})
.await
.expect("stream should start");
let mut got = Vec::new();
while let Some(item) = rx.recv().await {
got.push(item.unwrap());
}
assert_eq!(got, vec!["7:0", "7:1", "7:2"]);
}
#[tokio::test(start_paused = true)]
async fn streaming_stops_when_receiver_dropped() {
use std::sync::atomic::{AtomicUsize, Ordering as O};
struct InfiniteStream {
emitted: Arc<AtomicUsize>,
}
#[async_trait]
impl Dispatch for InfiniteStream {
async fn dispatch(&self, _p: &str, _v: Value) -> Result<Value, DispatchError> {
Err(DispatchError::new("unary not used"))
}
async fn dispatch_stream(
&self,
_p: &str,
_v: Value,
) -> Result<ChunkStream, DispatchError> {
let emitted = self.emitted.clone();
let s = futures::stream::unfold(0u64, move |n| {
let emitted = emitted.clone();
async move {
emitted.fetch_add(1, O::SeqCst);
tokio::task::yield_now().await;
Some((Ok(format!("chunk{n}")), n + 1))
}
});
Ok(Box::pin(s))
}
}
let emitted = Arc::new(AtomicUsize::new(0));
let sched = Scheduler::new(
GatewayConfig::default(),
Arc::new(FakeLimiter::new(100)),
Arc::new(InfiniteStream {
emitted: emitted.clone(),
}),
);
let mut rx = sched
.submit_stream(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({ "id": 1 }),
})
.await
.expect("stream should start");
assert!(rx.recv().await.is_some());
assert!(rx.recv().await.is_some());
drop(rx);
yield_many().await;
let a = emitted.load(O::SeqCst);
yield_many().await;
let b = emitted.load(O::SeqCst);
assert_eq!(a, b, "forwarding must stop once the client disconnects");
}
}