mod attempt;
pub mod auth;
mod budget;
pub mod http;
pub mod idempotency;
pub mod metrics;
pub mod quota;
#[cfg(feature = "redis-coordination")]
pub mod distributed;
use std::collections::{BTreeMap, BTreeSet, HashMap};
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<()>,
timing: Box<JobTiming>,
}
struct JobTiming {
enqueued_at: Instant,
deadline: Instant,
}
impl Job {
fn is_cancelled(&self) -> bool {
match &self.delivery {
Delivery::Unary(tx) => tx.is_closed(),
Delivery::Stream(tx) => tx.is_closed(),
}
}
fn send_timeout(self) {
let _ = self.started.send(());
match self.delivery {
Delivery::Unary(sender) => {
let _ = sender.send(Err(GatewayError::Timeout));
}
Delivery::Stream(sender) => {
let _ = sender.send(Err(GatewayError::Timeout));
}
}
}
}
#[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,
#[cfg(test)]
deadline_index_visits: AtomicU64,
}
#[derive(Default)]
struct QueueInner {
tiers: BTreeMap<Tier, BTreeSet<u64>>,
jobs: HashMap<u64, Job>,
deadlines: BTreeMap<Instant, BTreeSet<u64>>,
}
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,
#[cfg(test)]
deadline_index_visits: AtomicU64::new(0),
}
}
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 job_ids = inner.tiers.get(&tier).unwrap();
if let Some(job_id) = job_ids.first() {
let job = inner.jobs.get(job_id).unwrap();
let eff = self.effective_priority(tier, job.timing.enqueued_at, now);
let seqno = job.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 job_id = *inner.tiers.get(&tier).unwrap().first().unwrap();
remove_tier_index(inner, tier, job_id);
let job = inner.jobs.remove(&job_id).unwrap();
remove_deadline_index(inner, job.timing.deadline, job_id);
Some(job)
}
fn expire_and_earliest_deadline(&self, now: Instant) -> Option<Instant> {
let mut inner = self.inner.lock().unwrap();
loop {
let (&deadline, job_ids) = inner.deadlines.first_key_value()?;
let job_id = *job_ids.first().unwrap();
#[cfg(test)]
self.deadline_index_visits
.fetch_add(1, AtomicOrdering::Relaxed);
let cancelled = inner.jobs.get(&job_id).is_none_or(Job::is_cancelled);
if deadline > now && !cancelled {
return Some(deadline);
}
remove_deadline_index(&mut inner, deadline, job_id);
if let Some(job) = inner.jobs.remove(&job_id) {
remove_tier_index(&mut inner, job.key.tier, job_id);
if !cancelled {
job.send_timeout();
}
}
}
}
#[cfg(test)]
fn deadline_index_visits(&self) -> u64 {
self.deadline_index_visits.load(AtomicOrdering::Relaxed)
}
#[cfg(test)]
fn deadline_index_len(&self) -> usize {
self.inner
.lock()
.unwrap()
.deadlines
.values()
.map(BTreeSet::len)
.sum()
}
#[cfg(test)]
fn tier_index_len(&self) -> usize {
self.inner
.lock()
.unwrap()
.tiers
.values()
.map(BTreeSet::len)
.sum()
}
}
fn insert_deadline_index(inner: &mut QueueInner, deadline: Instant, job_id: u64) {
inner.deadlines.entry(deadline).or_default().insert(job_id);
}
fn remove_tier_index(inner: &mut QueueInner, tier: Tier, job_id: u64) {
if let Some(job_ids) = inner.tiers.get_mut(&tier) {
job_ids.remove(&job_id);
if job_ids.is_empty() {
inner.tiers.remove(&tier);
}
}
}
fn remove_deadline_index(inner: &mut QueueInner, deadline: Instant, job_id: u64) {
if let Some(job_ids) = inner.deadlines.get_mut(&deadline) {
job_ids.remove(&job_id);
if job_ids.is_empty() {
inner.deadlines.remove(&deadline);
}
}
}
#[async_trait]
impl RequestQueue for InMemoryQueue {
fn enqueue(&self, job: Job) -> Result<(), Job> {
{
let mut inner = self.inner.lock().unwrap();
if inner.jobs.len() >= self.max_depth {
return Err(job);
}
let job_id = job.key.seqno;
let tier = job.key.tier;
let deadline = job.timing.deadline;
inner.jobs.insert(job_id, job);
inner.tiers.entry(tier).or_default().insert(job_id);
insert_deadline_index(&mut inner, deadline, job_id);
}
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();
let job_id = job.key.seqno;
let tier = job.key.tier;
let deadline = job.timing.deadline;
inner.jobs.insert(job_id, job);
inner.tiers.entry(tier).or_default().insert(job_id);
insert_deadline_index(&mut inner, deadline, job_id);
}
async fn notified(&self) {
self.notify.notified().await;
}
fn depth(&self) -> usize {
self.inner.lock().unwrap().jobs.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 unary_job_timeout: Duration,
pub stream_job_timeout: Duration,
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,
unary_job_timeout: Duration::from_secs(2 * 60 * 60),
stream_job_timeout: Duration::from_secs(6 * 60 * 60),
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)
};
let duration_env = |key: &str, fallback: Duration| {
std::env::var(key)
.ok()
.and_then(|value| value.parse::<u64>().ok())
.filter(|milliseconds| *milliseconds > 0)
.map(Duration::from_millis)
.filter(|duration| Instant::now().checked_add(*duration).is_some())
.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,
),
unary_job_timeout: duration_env(
"LLMSHIM_GATEWAY_UNARY_JOB_TIMEOUT_MS",
d.unary_job_timeout,
),
stream_job_timeout: duration_env(
"LLMSHIM_GATEWAY_STREAM_JOB_TIMEOUT_MS",
d.stream_job_timeout,
),
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<InMemoryQueue>,
}
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, None).await
}
async fn submit_inner(
self: &Arc<Self>,
req: GatewayRequest,
policy_context: Option<crate::policy::DispatchPolicyContext>,
deadline_override: Option<Instant>,
) -> Result<Value, GatewayError> {
let deadline = deadline_override.unwrap_or_else(|| {
Instant::now()
.checked_add(self.config.unary_job_timeout)
.unwrap_or_else(|| Instant::now() + Duration::from_secs(2 * 60 * 60))
});
let policy_context = policy_context.map(|context| context.with_logical_deadline(deadline));
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,
timing: Box::new(JobTiming {
enqueued_at: Instant::now(),
deadline,
}),
};
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), None).await
}
pub(crate) async fn submit_with_policy_deadline(
self: &Arc<Self>,
req: GatewayRequest,
policy_context: crate::policy::DispatchPolicyContext,
deadline: Instant,
) -> Result<Value, GatewayError> {
self.submit_inner(req, Some(policy_context), Some(deadline))
.await
}
pub async fn submit_stream(
self: &Arc<Self>,
req: GatewayRequest,
) -> Result<mpsc::Receiver<StreamChunk>, GatewayError> {
self.submit_stream_inner(req, None, 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), None)
.await
}
pub(crate) async fn submit_stream_with_policy_deadline(
self: &Arc<Self>,
req: GatewayRequest,
policy_context: crate::policy::DispatchPolicyContext,
deadline: Instant,
) -> Result<mpsc::Receiver<StreamChunk>, GatewayError> {
self.submit_stream_inner(req, Some(policy_context), Some(deadline))
.await
}
async fn submit_stream_inner(
self: &Arc<Self>,
req: GatewayRequest,
policy_context: Option<crate::policy::DispatchPolicyContext>,
deadline_override: Option<Instant>,
) -> Result<mpsc::Receiver<StreamChunk>, GatewayError> {
let deadline = deadline_override.unwrap_or_else(|| {
Instant::now()
.checked_add(self.config.stream_job_timeout)
.unwrap_or_else(|| Instant::now() + Duration::from_secs(6 * 60 * 60))
});
let policy_context = policy_context.map(|context| context.with_logical_deadline(deadline));
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,
timing: Box::new(JobTiming {
enqueued_at: Instant::now(),
deadline,
}),
};
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<InMemoryQueue> {
let mut lanes = self.lanes.lock().unwrap();
if let Some(lane) = lanes.get(provider) {
return lane.queue.clone();
}
let queue = 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<InMemoryQueue>,
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 = tokio::select! {
biased;
_ = tokio::time::sleep_until(job.timing.deadline) => {
job.send_timeout();
continue;
}
result = sem.clone().acquire_owned() => match result {
Ok(permit) => permit,
Err(_) => break,
}
};
if job.is_cancelled() {
continue; }
let rate_admission = if policy_gated {
Ok(())
} else {
tokio::select! {
biased;
_ = tokio::time::sleep_until(job.timing.deadline) => {
job.send_timeout();
continue;
}
result = limiter.acquire(&key, job.permits) => result,
}
};
match rate_admission {
Ok(()) => {
if job.is_cancelled() {
continue; }
metrics::observe_ms(
metrics::QUEUE_WAIT,
&[("provider", &provider)],
job.timing.enqueued_at.elapsed().as_millis() as f64,
);
let Job {
payload,
policy_context,
delivery,
started,
timing,
..
} = job;
let deadline = timing.deadline;
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(mut tx) => {
let dispatch_future = match policy_context {
Some(context) => {
dispatch.dispatch_with_policy(&provider, payload, context)
}
None => dispatch.dispatch(&provider, payload),
};
tokio::pin!(dispatch_future);
let dispatch_result = tokio::select! {
biased;
_ = tx.closed() => return,
_ = tokio::time::sleep_until(deadline) => {
let _ = tx.send(Err(GatewayError::Timeout));
return;
}
result = &mut dispatch_future => result,
};
if Instant::now() >= deadline {
let _ = tx.send(Err(GatewayError::Timeout));
return;
}
match dispatch_result {
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(mut tx) => {
let open_future = match policy_context {
Some(context) => dispatch
.dispatch_stream_with_policy(&provider, payload, context),
None => dispatch.dispatch_stream(&provider, payload),
};
tokio::pin!(open_future);
let opened = tokio::select! {
biased;
_ = tx.closed() => return,
_ = tokio::time::sleep_until(deadline) => {
let _ = tx.send(Err(GatewayError::Timeout));
return;
}
result = &mut open_future => result,
};
if Instant::now() >= deadline {
let _ = tx.send(Err(GatewayError::Timeout));
return;
}
match opened {
Ok(mut upstream) => {
metrics::incr(metrics::DISPATCHED, plabels);
let (chunk_tx, chunk_rx) = mpsc::channel(17);
let terminal_permit =
match chunk_tx.clone().reserve_owned().await {
Ok(permit) => permit,
Err(_) => return,
};
if tx.send(Ok(chunk_rx)).is_err() {
return;
}
use futures::StreamExt;
let mut terminal_permit = Some(terminal_permit);
loop {
tokio::select! {
biased;
_ = chunk_tx.closed() => break,
_ = tokio::time::sleep_until(deadline) => {
if let Some(permit) = terminal_permit.take() {
permit.send(Err(GatewayError::Timeout));
}
break;
}
item = upstream.next() => match item {
Some(item) => {
let sent = tokio::select! {
biased;
_ = chunk_tx.closed() => false,
_ = tokio::time::sleep_until(deadline) => {
if let Some(permit) = terminal_permit.take() {
permit.send(Err(GatewayError::Timeout));
}
false
}
result = chunk_tx.send(item) => result.is_ok(),
};
if !sent {
break;
}
}
None => {
terminal_permit.take();
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);
let retry_deadline = Instant::now()
.checked_add(wait)
.unwrap_or_else(|| Instant::now() + Duration::from_secs(60));
let earliest_job_deadline = queue
.expire_and_earliest_deadline(Instant::now())
.unwrap_or(retry_deadline);
tokio::select! {
biased;
_ = tokio::time::sleep_until(earliest_job_deadline) => {}
_ = tokio::time::sleep_until(retry_deadline) => {}
_ = queue.notified() => {}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use std::collections::HashMap as Map;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
struct FakeLimiter {
permits: Mutex<Map<String, i64>>,
default: i64,
}
struct RetryHintLimiter(Duration);
#[async_trait]
impl RateLimiter for RetryHintLimiter {
async fn acquire(&self, _key: &RateKey, _permits: u32) -> Result<(), RetryAfter> {
Err(RetryAfter(self.0))
}
async fn penalize(&self, _key: &RateKey, _retry_after: Duration) {}
}
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)))
}
}
struct NoopAttemptPolicy;
impl crate::policy::AttemptPolicy for NoopAttemptPolicy {
fn acquire<'a>(
&'a self,
_attempt: &'a crate::policy::PreparedAttempt<'a>,
) -> crate::policy::AttemptPolicyFuture<'a, Result<(), crate::policy::AttemptPolicyRefusal>>
{
Box::pin(async { Ok(()) })
}
fn observe<'a>(
&'a self,
_attempt: &'a crate::policy::AttemptIdentity,
_event: crate::policy::AttemptEvent<'a>,
) -> crate::policy::AttemptPolicyFuture<'a, Result<(), crate::policy::AttemptPolicyError>>
{
Box::pin(async { Ok(()) })
}
fn observe_abandoned(
&self,
_attempt: &crate::policy::AttemptIdentity,
_outcome: crate::policy::AttemptOutcome,
) -> Result<(), crate::policy::AttemptPolicyError> {
Ok(())
}
}
struct LatchBlockedDispatch {
preparation_starts: Arc<std::sync::atomic::AtomicUsize>,
preparation_started: Arc<Notify>,
preparation_latch: Arc<Notify>,
}
struct PendingDispatch {
starts: Arc<AtomicUsize>,
dropped: Arc<AtomicBool>,
}
struct DispatchDropGuard(Arc<AtomicBool>);
impl Drop for DispatchDropGuard {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
#[async_trait]
impl Dispatch for PendingDispatch {
async fn dispatch(&self, _provider: &str, _payload: Value) -> Result<Value, DispatchError> {
self.starts.fetch_add(1, Ordering::SeqCst);
let _guard = DispatchDropGuard(self.dropped.clone());
futures::future::pending().await
}
async fn dispatch_stream(
&self,
_provider: &str,
_payload: Value,
) -> Result<ChunkStream, DispatchError> {
self.starts.fetch_add(1, Ordering::SeqCst);
let dropped = self.dropped.clone();
let stream = async_stream::stream! {
let _guard = DispatchDropGuard(dropped);
for index in 0..100_u32 {
yield Ok(index.to_string());
}
futures::future::pending::<()>().await;
};
Ok(Box::pin(stream))
}
}
#[async_trait]
impl Dispatch for LatchBlockedDispatch {
async fn dispatch(&self, _provider: &str, _payload: Value) -> Result<Value, DispatchError> {
unreachable!("policy-gated test uses dispatch_with_policy")
}
async fn dispatch_with_policy(
&self,
_provider: &str,
payload: Value,
_policy_context: crate::policy::DispatchPolicyContext,
) -> Result<Value, DispatchError> {
self.preparation_starts
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
self.preparation_started.notify_one();
self.preparation_latch.notified().await;
Ok(payload)
}
}
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,
timing: Box::new(JobTiming {
enqueued_at: Instant::now(),
deadline: Instant::now() + Duration::from_secs(60),
}),
},
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(start_paused = true)]
async fn deadline_index_work_is_constant_per_backoff_wake_and_drains_cleanly() {
let queue = InMemoryQueue::new(10_000);
let shared_deadline = Instant::now() + Duration::from_secs(60);
let mut receivers = Vec::new();
for sequence in 0..5_000_u64 {
let (mut job, receiver) = dummy_job(0, sequence);
job.timing.deadline = shared_deadline;
assert!(queue.enqueue(job).is_ok());
receivers.push(receiver);
}
for _ in 0..5_000 {
assert_eq!(
queue.expire_and_earliest_deadline(Instant::now()),
Some(shared_deadline)
);
}
assert_eq!(queue.deadline_index_visits(), 5_000);
assert_eq!(queue.deadline_index_len(), 5_000);
for _ in 0..5_000 {
let _ = queue.dequeue().await;
}
assert_eq!(queue.depth(), 0);
assert_eq!(queue.deadline_index_len(), 0);
drop(receivers);
}
#[tokio::test(start_paused = true)]
async fn repeated_expiry_behind_live_tier_head_keeps_every_index_bounded() {
let queue = InMemoryQueue::new(1_000);
let (mut head, head_receiver) = dummy_job(0, 0);
let head_deadline = Instant::now() + Duration::from_secs(1_000);
head.timing.deadline = head_deadline;
assert!(queue.enqueue(head).is_ok());
for wave in 0..20_u64 {
let deadline = Instant::now() + Duration::from_millis(1);
let mut receivers = Vec::new();
for offset in 0..100_u64 {
let (mut job, receiver) = dummy_job(0, wave * 100 + offset + 1);
job.timing.deadline = deadline;
assert!(queue.enqueue(job).is_ok());
receivers.push(receiver);
}
tokio::time::advance(Duration::from_millis(1)).await;
assert_eq!(
queue.expire_and_earliest_deadline(Instant::now()),
Some(head_deadline)
);
assert_eq!(queue.depth(), 1);
assert_eq!(queue.deadline_index_len(), 1);
assert_eq!(queue.tier_index_len(), 1);
for receiver in receivers {
assert!(matches!(
receiver.await.unwrap(),
Err(GatewayError::Timeout)
));
}
}
let _ = queue.dequeue().await;
assert_eq!(queue.depth(), 0);
assert_eq!(queue.deadline_index_len(), 0);
assert_eq!(queue.tier_index_len(), 0);
drop(head_receiver);
}
#[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]
async fn policy_dispatch_preparation_is_bounded_by_scheduler_concurrency() {
let preparation_starts = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let preparation_started = Arc::new(Notify::new());
let first_preparation_started = preparation_started.notified();
let preparation_latch = Arc::new(Notify::new());
let scheduler = Scheduler::new(
GatewayConfig {
max_concurrency_per_provider: 1,
..Default::default()
},
Arc::new(FakeLimiter::new(100)),
Arc::new(LatchBlockedDispatch {
preparation_starts: preparation_starts.clone(),
preparation_started: preparation_started.clone(),
preparation_latch: preparation_latch.clone(),
}),
);
let policy_context = crate::policy::DispatchPolicyContext::new(Arc::new(NoopAttemptPolicy));
let mut request_handles = Vec::new();
for id in 0..3 {
let scheduler = scheduler.clone();
let policy_context = policy_context.clone();
request_handles.push(tokio::spawn(async move {
scheduler
.submit_with_policy(
GatewayRequest {
provider: "blocked".into(),
tier: 0,
permits: 1,
payload: json!({"id": id}),
},
policy_context,
)
.await
}));
}
tokio::time::timeout(Duration::from_secs(1), first_preparation_started)
.await
.expect("first request should enter dispatch preparation");
assert_eq!(
preparation_starts.load(std::sync::atomic::Ordering::SeqCst),
1
);
for request_handle in request_handles {
request_handle.abort();
}
preparation_latch.notify_waiters();
}
#[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 unary_deadline_interrupts_long_rate_limit_backoff() {
let order = Arc::new(Mutex::new(Vec::new()));
let scheduler = Scheduler::new(
GatewayConfig {
max_wait: Duration::from_millis(500),
unary_job_timeout: Duration::from_millis(20),
..GatewayConfig::default()
},
Arc::new(RetryHintLimiter(Duration::from_millis(200))),
Arc::new(RecordingDispatch {
order: order.clone(),
}),
);
let submitted_scheduler = scheduler.clone();
let handle = tokio::spawn(async move {
submitted_scheduler
.submit(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({"id": 1}),
})
.await
});
yield_many().await;
tokio::time::advance(Duration::from_millis(20)).await;
yield_many().await;
assert!(matches!(handle.await.unwrap(), Err(GatewayError::Timeout)));
assert!(order.lock().unwrap().is_empty());
assert_eq!(scheduler.queue_depth("p"), 0);
}
#[tokio::test(start_paused = true)]
async fn stream_deadline_interrupts_long_rate_limit_backoff() {
let order = Arc::new(Mutex::new(Vec::new()));
let scheduler = Scheduler::new(
GatewayConfig {
max_wait: Duration::from_millis(500),
stream_job_timeout: Duration::from_millis(20),
..GatewayConfig::default()
},
Arc::new(RetryHintLimiter(Duration::from_millis(200))),
Arc::new(RecordingDispatch {
order: order.clone(),
}),
);
let submitted_scheduler = scheduler.clone();
let handle = tokio::spawn(async move {
submitted_scheduler
.submit_stream(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({"id": 1}),
})
.await
});
yield_many().await;
tokio::time::advance(Duration::from_millis(20)).await;
yield_many().await;
assert!(matches!(handle.await.unwrap(), Err(GatewayError::Timeout)));
assert!(order.lock().unwrap().is_empty());
assert_eq!(scheduler.queue_depth("p"), 0);
}
#[tokio::test(start_paused = true)]
async fn backoff_wakes_for_earliest_deadline_across_priority_tiers() {
let order = Arc::new(Mutex::new(Vec::new()));
let scheduler = Scheduler::new(
GatewayConfig {
max_wait: Duration::from_secs(1),
..GatewayConfig::default()
},
Arc::new(RetryHintLimiter(Duration::from_millis(200))),
Arc::new(RecordingDispatch {
order: order.clone(),
}),
);
let now = Instant::now();
let high_scheduler = scheduler.clone();
let high = tokio::spawn(async move {
high_scheduler
.submit_inner(
GatewayRequest {
provider: "p".into(),
tier: 10,
permits: 1,
payload: json!({"id": 10}),
},
None,
Some(now + Duration::from_millis(200)),
)
.await
});
yield_many().await;
let low_scheduler = scheduler.clone();
let low = tokio::spawn(async move {
low_scheduler
.submit_inner(
GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({"id": 1}),
},
None,
Some(now + Duration::from_millis(20)),
)
.await
});
yield_many().await;
tokio::time::advance(Duration::from_millis(20)).await;
yield_many().await;
assert!(matches!(low.await.unwrap(), Err(GatewayError::Timeout)));
assert!(order.lock().unwrap().is_empty());
high.abort();
tokio::time::advance(Duration::from_millis(180)).await;
yield_many().await;
assert_eq!(scheduler.queue_depth("p"), 0);
}
#[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 dropping_unary_submitter_cancels_dispatch_and_releases_capacity() {
let starts = Arc::new(AtomicUsize::new(0));
let dropped = Arc::new(AtomicBool::new(false));
let config = GatewayConfig {
max_concurrency_per_provider: 1,
unary_job_timeout: Duration::from_secs(30),
..GatewayConfig::default()
};
let scheduler = Scheduler::new(
config,
Arc::new(FakeLimiter::new(100)),
Arc::new(PendingDispatch {
starts: starts.clone(),
dropped: dropped.clone(),
}),
);
let first_scheduler = scheduler.clone();
let first = tokio::spawn(async move {
first_scheduler
.submit(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({}),
})
.await
});
while starts.load(Ordering::SeqCst) == 0 {
tokio::task::yield_now().await;
}
first.abort();
yield_many().await;
assert!(dropped.load(Ordering::SeqCst));
let second_scheduler = scheduler.clone();
let second = tokio::spawn(async move {
second_scheduler
.submit(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({}),
})
.await
});
while starts.load(Ordering::SeqCst) < 2 {
tokio::task::yield_now().await;
}
second.abort();
}
#[tokio::test(start_paused = true)]
async fn unpolled_stream_gets_terminal_timeout_and_releases_capacity() {
let starts = Arc::new(AtomicUsize::new(0));
let dropped = Arc::new(AtomicBool::new(false));
let config = GatewayConfig {
max_concurrency_per_provider: 1,
stream_job_timeout: Duration::from_secs(2),
..GatewayConfig::default()
};
let scheduler = Scheduler::new(
config,
Arc::new(FakeLimiter::new(100)),
Arc::new(PendingDispatch {
starts: starts.clone(),
dropped: dropped.clone(),
}),
);
let mut first = scheduler
.submit_stream(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({}),
})
.await
.unwrap();
yield_many().await;
tokio::time::advance(Duration::from_secs(2)).await;
yield_many().await;
assert!(dropped.load(Ordering::SeqCst));
let mut saw_timeout = false;
while let Some(item) = first.recv().await {
if matches!(item, Err(GatewayError::Timeout)) {
saw_timeout = true;
break;
}
}
assert!(saw_timeout, "deadline must not become clean EOF");
let second = scheduler
.submit_stream(GatewayRequest {
provider: "p".into(),
tier: 0,
permits: 1,
payload: json!({}),
})
.await
.unwrap();
assert_eq!(starts.load(Ordering::SeqCst), 2);
drop(second);
}
#[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");
}
#[tokio::test]
async fn quiet_stream_releases_capacity_when_receiver_drops() {
struct QuietStreamDispatch {
unary_dispatches: Arc<std::sync::atomic::AtomicUsize>,
}
#[async_trait]
impl Dispatch for QuietStreamDispatch {
async fn dispatch(
&self,
_provider: &str,
payload: Value,
) -> Result<Value, DispatchError> {
self.unary_dispatches
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(payload)
}
async fn dispatch_stream(
&self,
_provider: &str,
_payload: Value,
) -> Result<ChunkStream, DispatchError> {
Ok(Box::pin(futures::stream::pending()))
}
}
let unary_dispatches = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let scheduler = Scheduler::new(
GatewayConfig {
max_concurrency_per_provider: 1,
..Default::default()
},
Arc::new(FakeLimiter::new(100)),
Arc::new(QuietStreamDispatch {
unary_dispatches: unary_dispatches.clone(),
}),
);
let stream_receiver = scheduler
.submit_stream(GatewayRequest {
provider: "quiet".into(),
tier: 0,
permits: 1,
payload: json!({}),
})
.await
.expect("quiet stream should open");
drop(stream_receiver);
let response = tokio::time::timeout(
Duration::from_secs(1),
scheduler.submit(GatewayRequest {
provider: "quiet".into(),
tier: 0,
permits: 1,
payload: json!({"id": 2}),
}),
)
.await
.expect("receiver close should promptly release capacity")
.expect("second request should dispatch");
assert_eq!(response["id"], 2);
assert_eq!(
unary_dispatches.load(std::sync::atomic::Ordering::SeqCst),
1
);
}
}