use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::task::{Context, Poll};
use std::time::{Duration, Instant, SystemTime};
use dashmap::DashMap;
use tower::{Layer, Service};
use super::cost::observe_stream_usage;
use super::types::{LlmRequest, LlmResponse};
use crate::client::BoxFuture;
use crate::cost;
use crate::error::{LiterLlmError, Result};
use crate::types::Usage;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct RateLimitConfig {
pub rpm: Option<u32>,
pub tpm: Option<u64>,
pub window: Duration,
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
rpm: None,
tpm: None,
window: Duration::from_secs(60),
}
}
}
struct ModelRateState {
request_count: u64,
token_count: u64,
window_start: Instant,
}
impl ModelRateState {
fn new() -> Self {
Self {
request_count: 0,
token_count: 0,
window_start: Instant::now(),
}
}
fn maybe_reset(&mut self, window: Duration) {
if self.window_start.elapsed() >= window {
self.request_count = 0;
self.token_count = 0;
self.window_start = Instant::now();
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct ModelRateLimitLayer {
config: RateLimitConfig,
state: Arc<DashMap<String, ModelRateState>>,
}
impl ModelRateLimitLayer {
#[must_use]
pub fn new(config: RateLimitConfig) -> Self {
Self {
config,
state: Arc::new(DashMap::new()),
}
}
}
impl<S> Layer<S> for ModelRateLimitLayer {
type Service = ModelRateLimitService<S>;
fn layer(&self, inner: S) -> Self::Service {
ModelRateLimitService {
inner,
config: self.config.clone(),
state: Arc::clone(&self.state),
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct ModelRateLimitService<S> {
inner: S,
config: RateLimitConfig,
state: Arc<DashMap<String, ModelRateState>>,
}
impl<S: Clone> Clone for ModelRateLimitService<S> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
config: self.config.clone(),
state: Arc::clone(&self.state),
}
}
}
impl<S> Service<LlmRequest> for ModelRateLimitService<S>
where
S: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Send + 'static,
S::Future: Send + 'static,
{
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<()>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: LlmRequest) -> Self::Future {
let model = req.model().unwrap_or("unknown").to_owned();
let config = self.config.clone();
let state = Arc::clone(&self.state);
{
let mut entry = state.entry(model.clone()).or_insert_with(ModelRateState::new);
entry.maybe_reset(config.window);
if let Some(rpm) = config.rpm
&& entry.request_count >= u64::from(rpm)
{
return Box::pin(async move {
Err(LiterLlmError::RateLimited {
message: format!(
"model {model} exceeded {rpm} requests per {:.0}s window",
config.window.as_secs_f64()
),
retry_after: Some(config.window),
})
});
}
if let Some(tpm) = config.tpm
&& entry.token_count >= tpm
{
return Box::pin(async move {
Err(LiterLlmError::RateLimited {
message: format!(
"model {model} exceeded {tpm} tokens per {:.0}s window",
config.window.as_secs_f64()
),
retry_after: Some(config.window),
})
});
}
entry.request_count += 1;
}
let fut = self.inner.call(req);
Box::pin(async move {
let resp = fut.await?;
match resp {
LlmResponse::ChatStream(stream) => {
let model_for_completion = model.clone();
let state_for_completion = Arc::clone(&state);
let config_for_completion = config.clone();
let wrapped = observe_stream_usage(stream, move |usage| {
record_tokens(
&state_for_completion,
&config_for_completion,
&model_for_completion,
usage.as_ref(),
);
});
Ok(LlmResponse::ChatStream(wrapped))
}
other => {
record_tokens(&state, &config, &model, other.usage());
Ok(other)
}
}
})
}
}
fn record_tokens(
state: &DashMap<String, ModelRateState>,
config: &RateLimitConfig,
model: &str,
usage: Option<&Usage>,
) {
let Some(usage) = usage else { return };
let total_tokens = usage.prompt_tokens + usage.completion_tokens;
if let Some(mut entry) = state.get_mut(model) {
entry.maybe_reset(config.window);
entry.token_count += total_tokens;
}
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub struct CostRateLimitConfig {
pub max_usd_per_minute: Option<f64>,
pub max_usd_per_hour: Option<f64>,
pub max_usd_per_day: Option<f64>,
}
#[derive(Debug)]
struct CostWindow {
spend_mc: AtomicU64,
window_start_secs: AtomicU64,
window_secs: u64,
}
impl CostWindow {
fn new(window: Duration) -> Self {
let now = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
Self {
spend_mc: AtomicU64::new(0),
window_start_secs: AtomicU64::new(now),
window_secs: window.as_secs(),
}
}
fn spend_usd(&self, now_secs: u64) -> f64 {
let start = self.window_start_secs.load(Ordering::Acquire);
if now_secs.saturating_sub(start) >= self.window_secs {
let old_mc = self.spend_mc.load(Ordering::Acquire);
if self
.window_start_secs
.compare_exchange(start, now_secs, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
self.spend_mc.fetch_sub(old_mc, Ordering::AcqRel);
}
}
self.spend_mc.load(Ordering::Acquire) as f64 / 1_000_000.0
}
fn add(&self, usd: f64, now_secs: u64) {
let _ = self.spend_usd(now_secs);
if usd > 0.0 {
let mc = (usd * 1_000_000.0).round() as u64;
self.spend_mc.fetch_add(mc, Ordering::AcqRel);
}
}
}
#[derive(Debug)]
struct CostRateLimitState {
per_minute: CostWindow,
per_hour: CostWindow,
per_day: CostWindow,
}
impl CostRateLimitState {
fn new() -> Self {
Self {
per_minute: CostWindow::new(Duration::from_secs(60)),
per_hour: CostWindow::new(Duration::from_secs(3600)),
per_day: CostWindow::new(Duration::from_secs(86_400)),
}
}
fn now_secs() -> u64 {
SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
fn check(&self, config: &CostRateLimitConfig) -> Option<LiterLlmError> {
let now = Self::now_secs();
if let Some(limit) = config.max_usd_per_minute {
let spend = self.per_minute.spend_usd(now);
if spend >= limit {
return Some(LiterLlmError::RateLimited {
message: format!("cost rate limit exceeded: ${spend:.6} >= ${limit:.6} per minute"),
retry_after: Some(Duration::from_secs(60)),
});
}
}
if let Some(limit) = config.max_usd_per_hour {
let spend = self.per_hour.spend_usd(now);
if spend >= limit {
return Some(LiterLlmError::RateLimited {
message: format!("cost rate limit exceeded: ${spend:.6} >= ${limit:.6} per hour"),
retry_after: Some(Duration::from_secs(3600)),
});
}
}
if let Some(limit) = config.max_usd_per_day {
let spend = self.per_day.spend_usd(now);
if spend >= limit {
return Some(LiterLlmError::RateLimited {
message: format!("cost rate limit exceeded: ${spend:.6} >= ${limit:.6} per day"),
retry_after: Some(Duration::from_secs(86_400)),
});
}
}
None
}
fn record(&self, usd: f64) {
let now = Self::now_secs();
self.per_minute.add(usd, now);
self.per_hour.add(usd, now);
self.per_day.add(usd, now);
}
}
#[cfg_attr(alef, alef(skip))]
pub struct CostRateLimitLayer {
config: CostRateLimitConfig,
state: Arc<CostRateLimitState>,
}
impl CostRateLimitLayer {
#[must_use]
pub fn new(config: CostRateLimitConfig) -> Self {
Self {
config,
state: Arc::new(CostRateLimitState::new()),
}
}
}
impl<S> Layer<S> for CostRateLimitLayer {
type Service = CostRateLimitService<S>;
fn layer(&self, inner: S) -> Self::Service {
CostRateLimitService {
inner,
config: self.config.clone(),
state: Arc::clone(&self.state),
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct CostRateLimitService<S> {
inner: S,
config: CostRateLimitConfig,
state: Arc<CostRateLimitState>,
}
impl<S: Clone> Clone for CostRateLimitService<S> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
config: self.config.clone(),
state: Arc::clone(&self.state),
}
}
}
impl<S> Service<LlmRequest> for CostRateLimitService<S>
where
S: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Send + 'static,
S::Future: Send + 'static,
{
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<()>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: LlmRequest) -> Self::Future {
let model = req.model().unwrap_or("unknown").to_owned();
let config = self.config.clone();
let state = Arc::clone(&self.state);
if let Some(err) = state.check(&config) {
return Box::pin(async move { Err(err) });
}
let fut = self.inner.call(req);
Box::pin(async move {
let resp = fut.await?;
match resp {
LlmResponse::ChatStream(stream) => {
let model_for_completion = model.clone();
let state_for_completion = Arc::clone(&state);
let wrapped = observe_stream_usage(stream, move |usage| {
record_cost_window(&state_for_completion, &model_for_completion, usage.as_ref());
});
Ok(LlmResponse::ChatStream(wrapped))
}
other => {
record_cost_window(&state, &model, other.usage());
Ok(other)
}
}
})
}
}
fn record_cost_window(state: &CostRateLimitState, model: &str, usage: Option<&Usage>) {
let Some(usage) = usage else { return };
if let Some(usd) = cost::completion_cost(model, usage.prompt_tokens, usage.completion_tokens) {
state.record(usd);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tower::tests_common::{MockClient, chat_req};
use crate::tower::service::LlmService;
use crate::tower::types::LlmRequest;
#[tokio::test]
async fn allows_requests_under_rpm_limit() {
let config = RateLimitConfig {
rpm: Some(5),
tpm: None,
window: Duration::from_secs(60),
};
let layer = ModelRateLimitLayer::new(config);
let inner = LlmService::new(MockClient::ok());
let mut svc = layer.layer(inner);
for _ in 0..5 {
let resp = svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok(), "requests under limit should succeed");
}
}
#[tokio::test]
async fn rejects_requests_over_rpm_limit() {
let config = RateLimitConfig {
rpm: Some(2),
tpm: None,
window: Duration::from_secs(60),
};
let layer = ModelRateLimitLayer::new(config);
let inner = LlmService::new(MockClient::ok());
let mut svc = layer.layer(inner);
svc.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("service call should not fail");
svc.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("service call should not fail");
let err = svc
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect_err("should be rate limited");
assert!(matches!(err, LiterLlmError::RateLimited { .. }));
}
#[tokio::test]
async fn independent_models_have_separate_limits() {
let config = RateLimitConfig {
rpm: Some(1),
tpm: None,
window: Duration::from_secs(60),
};
let layer = ModelRateLimitLayer::new(config);
let inner = LlmService::new(MockClient::ok());
let mut svc = layer.layer(inner);
svc.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("service call should not fail");
svc.call(LlmRequest::Chat(chat_req("gpt-3.5-turbo")))
.await
.expect("service call should not fail");
}
#[tokio::test]
async fn tpm_limit_rejects_after_threshold() {
let config = RateLimitConfig {
rpm: None,
tpm: Some(10),
window: Duration::from_secs(60),
};
let layer = ModelRateLimitLayer::new(config);
let inner = LlmService::new(MockClient::ok());
let mut svc = layer.layer(inner);
svc.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("service call should not fail");
let err = svc
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect_err("should be rate limited by TPM");
assert!(matches!(err, LiterLlmError::RateLimited { .. }));
}
#[tokio::test]
async fn unlimited_config_allows_all_requests() {
let config = RateLimitConfig::default();
let layer = ModelRateLimitLayer::new(config);
let inner = LlmService::new(MockClient::ok());
let mut svc = layer.layer(inner);
for _ in 0..100 {
assert!(svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await.is_ok());
}
}
#[tokio::test]
async fn cost_rate_limit_rejects_when_projected_exceeds_max() {
let config = CostRateLimitConfig {
max_usd_per_minute: Some(0.01),
max_usd_per_hour: None,
max_usd_per_day: None,
};
let layer = CostRateLimitLayer::new(config);
let inner = LlmService::new(MockClient::ok());
let mut svc = layer.layer(inner);
let config2 = CostRateLimitConfig {
max_usd_per_minute: Some(0.000001),
max_usd_per_hour: None,
max_usd_per_day: None,
};
let layer2 = CostRateLimitLayer::new(config2);
let inner2 = LlmService::new(MockClient::ok());
let mut svc2 = layer2.layer(inner2);
svc2.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("first call should succeed");
let err = svc2
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect_err("should be rate limited by cost");
assert!(
matches!(err, LiterLlmError::RateLimited { .. }),
"expected RateLimited, got {err:?}"
);
svc.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("request under cost limit should succeed");
}
#[tokio::test]
async fn cost_rate_limit_unlimited_config_allows_all_requests() {
let config = CostRateLimitConfig::default();
let layer = CostRateLimitLayer::new(config);
let inner = LlmService::new(MockClient::ok());
let mut svc = layer.layer(inner);
for _ in 0..20 {
assert!(svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await.is_ok());
}
}
#[tokio::test]
async fn cost_rate_limit_propagates_inner_errors() {
let config = CostRateLimitConfig {
max_usd_per_minute: Some(100.0),
max_usd_per_hour: None,
max_usd_per_day: None,
};
let layer = CostRateLimitLayer::new(config);
let inner = LlmService::new(MockClient::failing_timeout());
let mut svc = layer.layer(inner);
let err = svc
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect_err("inner error should propagate");
assert!(matches!(err, LiterLlmError::Timeout));
}
#[test]
fn cost_window_rollover_under_concurrent_threads_does_not_undercount() {
use std::sync::Barrier;
use std::thread;
let window = Arc::new(CostWindow::new(Duration::from_secs(1)));
let future_now = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
+ 2;
const WRITERS: usize = 200;
let barrier = Arc::new(Barrier::new(WRITERS));
let mut handles = Vec::with_capacity(WRITERS);
for _ in 0..WRITERS {
let w = Arc::clone(&window);
let b = Arc::clone(&barrier);
handles.push(thread::spawn(move || {
b.wait();
w.add(0.10, future_now);
}));
}
for h in handles {
h.join().expect("writer must not panic");
}
let total = window.spend_mc.load(Ordering::Relaxed) as f64 / 1_000_000.0;
assert!(
(total - 20.0_f64).abs() < 1e-4,
"expected $20.00 total after 200 concurrent adds at rollover; got ${total:.6}"
);
}
#[tokio::test]
async fn model_rate_limit_records_tokens_for_streamed_response() {
use std::collections::VecDeque;
use futures_core::Stream;
use futures_util::StreamExt as _;
use crate::client::BoxStream;
use crate::types::ChatCompletionChunk;
struct ChunkStream(VecDeque<ChatCompletionChunk>);
impl Stream for ChunkStream {
type Item = Result<ChatCompletionChunk>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
std::task::Poll::Ready(self.0.pop_front().map(Ok))
}
}
fn usage_chunk(usage: Option<Usage>) -> ChatCompletionChunk {
ChatCompletionChunk {
id: "chunk".into(),
object: "chat.completion.chunk".into(),
created: 0,
model: "gpt-4".into(),
choices: vec![],
usage,
system_fingerprint: None,
service_tier: None,
}
}
#[derive(Clone)]
struct StreamingUsageService;
impl tower::Service<LlmRequest> for StreamingUsageService {
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, _cx: &mut std::task::Context<'_>) -> std::task::Poll<Result<()>> {
std::task::Poll::Ready(Ok(()))
}
fn call(&mut self, _req: LlmRequest) -> Self::Future {
Box::pin(async move {
let usage = Usage {
prompt_tokens: 30,
completion_tokens: 20,
total_tokens: 50,
prompt_tokens_details: None,
};
let chunks = VecDeque::from([usage_chunk(None), usage_chunk(Some(usage))]);
let stream: BoxStream<'static, Result<ChatCompletionChunk>> = Box::pin(ChunkStream(chunks));
Ok(LlmResponse::ChatStream(stream))
})
}
}
let config = RateLimitConfig {
rpm: None,
tpm: Some(50),
window: Duration::from_secs(60),
};
let layer = ModelRateLimitLayer::new(config);
let mut svc = layer.layer(StreamingUsageService);
let resp = svc
.call(LlmRequest::ChatStream(chat_req("gpt-4")))
.await
.expect("streamed call should succeed");
let LlmResponse::ChatStream(mut stream) = resp else {
panic!("expected a ChatStream response");
};
while stream.next().await.is_some() {}
let err = svc
.call(LlmRequest::ChatStream(chat_req("gpt-4")))
.await
.expect_err("second call must be rejected once the streamed call's 50 tokens are recorded");
assert!(
matches!(err, LiterLlmError::RateLimited { .. }),
"expected RateLimited once streamed tokens push the window to its TPM limit; got {err:?}"
);
}
}