#[cfg(test)]
use crate::frontend::generation::GENERATION_RETRY_AFTER_SECS;
#[cfg(test)]
use axum::http::StatusCode;
use openai_frontend::OpenAiError;
#[cfg(test)]
use openai_frontend::OpenAiErrorKind;
use openai_frontend::OpenAiResult;
use std::sync::Arc;
use std::sync::Mutex;
use tokio::sync::Notify;
#[cfg(test)]
use std::sync::Condvar;
#[cfg(test)]
use std::time::Duration;
#[cfg(test)]
use std::time::Instant;
pub const DECODE_BATCH_HEADROOM_TOKENS: usize = 512;
#[derive(Debug)]
pub(super) struct GenerationTokenBudget {
capacity_tokens: usize,
state: Mutex<GenerationTokenBudgetState>,
#[cfg(test)]
released: Condvar,
released_async: Notify,
}
#[derive(Debug, Default)]
struct GenerationTokenBudgetState {
active_tokens: usize,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) struct GenerationTokenBudgetRequest {
prompt_tokens: usize,
max_tokens: u32,
}
#[derive(Debug)]
pub(super) struct GenerationTokenReservation {
budget: Arc<GenerationTokenBudget>,
tokens: usize,
active_tokens_after_reservation: usize,
}
impl GenerationTokenBudget {
pub(super) fn new(ctx_size: usize) -> Self {
Self {
capacity_tokens: ctx_size.max(1),
state: Mutex::new(GenerationTokenBudgetState::default()),
#[cfg(test)]
released: Condvar::new(),
released_async: Notify::new(),
}
}
pub(super) fn try_reserve(
self: &Arc<Self>,
request: GenerationTokenBudgetRequest,
) -> OpenAiResult<Option<GenerationTokenReservation>> {
let tokens = request.reservation_tokens();
self.ensure_request_fits(tokens)?;
let mut state = self
.state
.lock()
.map_err(|_| OpenAiError::backend("generation token budget lock poisoned"))?;
if state.active_tokens.saturating_add(tokens) > self.capacity_tokens {
return Ok(None);
}
state.active_tokens = state.active_tokens.saturating_add(tokens);
Ok(Some(GenerationTokenReservation {
budget: self.clone(),
tokens,
active_tokens_after_reservation: state.active_tokens,
}))
}
pub(super) fn can_reserve_now(
&self,
request: GenerationTokenBudgetRequest,
) -> OpenAiResult<bool> {
let tokens = request.reservation_tokens();
self.ensure_request_fits(tokens)?;
let state = self
.state
.lock()
.map_err(|_| OpenAiError::backend("generation token budget lock poisoned"))?;
Ok(state.active_tokens.saturating_add(tokens) <= self.capacity_tokens)
}
pub(super) fn release_notification(&self) -> &Notify {
&self.released_async
}
#[cfg(test)]
pub(super) fn reserve(
self: &Arc<Self>,
request: GenerationTokenBudgetRequest,
admission_timeout: Duration,
) -> OpenAiResult<GenerationTokenReservation> {
self.reserve_cancellable(request, admission_timeout, None)
}
#[cfg(test)]
fn reserve_cancellable(
self: &Arc<Self>,
request: GenerationTokenBudgetRequest,
admission_timeout: Duration,
cancellation: Option<&openai_frontend::CancellationToken>,
) -> OpenAiResult<GenerationTokenReservation> {
let tokens = request.reservation_tokens();
self.ensure_request_fits(tokens)?;
let deadline = Instant::now() + admission_timeout;
let mut state = self
.state
.lock()
.map_err(|_| OpenAiError::backend("generation token budget lock poisoned"))?;
loop {
if cancellation.is_some_and(openai_frontend::CancellationToken::is_cancelled) {
return Err(OpenAiError::backend("request cancelled"));
}
if state.active_tokens.saturating_add(tokens) <= self.capacity_tokens {
state.active_tokens = state.active_tokens.saturating_add(tokens);
return Ok(GenerationTokenReservation {
budget: self.clone(),
tokens,
active_tokens_after_reservation: state.active_tokens,
});
}
let now = Instant::now();
if now >= deadline {
return Err(generation_token_budget_timeout_error(
admission_timeout,
tokens,
state.active_tokens,
self.capacity_tokens,
));
}
let wait_for = deadline.saturating_duration_since(now).min(
cancellation
.map(|_| Duration::from_millis(10))
.unwrap_or(admission_timeout),
);
let (next_state, wait_result) = self
.released
.wait_timeout(state, wait_for)
.map_err(|_| OpenAiError::backend("generation token budget lock poisoned"))?;
state = next_state;
if wait_result.timed_out() && Instant::now() >= deadline {
return Err(generation_token_budget_timeout_error(
admission_timeout,
tokens,
state.active_tokens,
self.capacity_tokens,
));
}
}
}
pub(super) fn capacity_tokens(&self) -> usize {
self.capacity_tokens
}
fn ensure_request_fits(&self, requested_tokens: usize) -> OpenAiResult<()> {
if requested_tokens <= self.capacity_tokens {
return Ok(());
}
Err(OpenAiError::context_length_exceeded(format!(
"request requires {requested_tokens} KV tokens but the runtime pool holds {}",
self.capacity_tokens
)))
}
#[cfg(test)]
pub(super) fn active_tokens(&self) -> usize {
self.state
.lock()
.expect("generation token budget lock")
.active_tokens
}
}
impl GenerationTokenBudgetRequest {
pub(super) fn new(prompt_tokens: usize, max_tokens: u32) -> Self {
Self {
prompt_tokens,
max_tokens,
}
}
pub(super) fn reservation_tokens(self) -> usize {
let max_tokens = usize::try_from(self.max_tokens).unwrap_or(usize::MAX);
let decode_headroom = max_tokens.min(DECODE_BATCH_HEADROOM_TOKENS);
self.prompt_tokens.saturating_add(decode_headroom)
}
}
impl GenerationTokenReservation {
pub(super) fn tokens(&self) -> usize {
self.tokens
}
pub(super) fn active_tokens_after_reservation(&self) -> usize {
self.active_tokens_after_reservation
}
}
impl Drop for GenerationTokenReservation {
fn drop(&mut self) {
if self.tokens == 0 {
return;
}
if let Ok(mut state) = self.budget.state.lock() {
state.active_tokens = state.active_tokens.saturating_sub(self.tokens);
#[cfg(test)]
self.budget.released.notify_one();
self.budget.released_async.notify_waiters();
}
}
}
#[cfg(test)]
fn generation_token_budget_timeout_error(
timeout: Duration,
requested_tokens: usize,
active_tokens: usize,
capacity_tokens: usize,
) -> OpenAiError {
OpenAiError::from_kind(
StatusCode::TOO_MANY_REQUESTS,
OpenAiErrorKind::RateLimit,
format!(
"timed out waiting for KV token budget after {} seconds \
(requested_tokens={requested_tokens}, active_tokens={active_tokens}, \
capacity_tokens={capacity_tokens})",
timeout.as_secs()
),
)
.with_retry_after_secs(GENERATION_RETRY_AFTER_SECS)
}
#[cfg(test)]
mod tests {
use super::*;
use std::{
sync::mpsc,
thread,
time::{Duration, Instant},
};
#[test]
fn generation_token_budget_reserves_and_releases_tokens() {
let budget = Arc::new(GenerationTokenBudget::new(1_024));
let reservation = budget
.reserve(
GenerationTokenBudgetRequest::new(400, 128),
Duration::from_millis(10),
)
.unwrap();
assert_eq!(reservation.tokens(), 528);
assert_eq!(reservation.active_tokens_after_reservation(), 528);
assert_eq!(budget.active_tokens(), 528);
drop(reservation);
assert_eq!(budget.active_tokens(), 0);
}
#[test]
fn generation_token_budget_clamps_decode_headroom() {
let budget = Arc::new(GenerationTokenBudget::new(8_192));
let reservation = budget
.reserve(
GenerationTokenBudgetRequest::new(4_506, 4_096),
Duration::from_millis(10),
)
.unwrap();
assert_eq!(reservation.tokens(), 5_018);
assert_eq!(reservation.active_tokens_after_reservation(), 5_018);
}
#[test]
fn generation_token_budget_waits_for_tokens_to_release() {
let budget = Arc::new(GenerationTokenBudget::new(1_000));
let first = budget
.reserve(
GenerationTokenBudgetRequest::new(800, 0),
Duration::from_millis(10),
)
.unwrap();
let waiter = budget.clone();
let (tx, rx) = mpsc::channel();
let handle = thread::spawn(move || {
let reservation = waiter
.reserve(
GenerationTokenBudgetRequest::new(300, 0),
Duration::from_secs(1),
)
.unwrap();
tx.send(reservation.active_tokens_after_reservation())
.unwrap();
});
assert!(rx.recv_timeout(Duration::from_millis(20)).is_err());
drop(first);
assert_eq!(rx.recv_timeout(Duration::from_secs(1)).unwrap(), 300);
handle.join().unwrap();
assert_eq!(budget.active_tokens(), 0);
}
#[test]
fn generation_token_budget_times_out_when_tokens_do_not_fit() {
let budget = Arc::new(GenerationTokenBudget::new(1_000));
let _first = budget
.reserve(
GenerationTokenBudgetRequest::new(800, 0),
Duration::from_millis(10),
)
.unwrap();
let started = Instant::now();
let error = budget
.reserve(
GenerationTokenBudgetRequest::new(300, 0),
Duration::from_millis(15),
)
.unwrap_err();
assert!(started.elapsed() >= Duration::from_millis(10));
assert_eq!(error.status(), StatusCode::TOO_MANY_REQUESTS);
assert_eq!(
error.body().error.code.as_deref(),
Some("rate_limit_exceeded")
);
assert!(error.to_string().contains("KV token budget"));
assert_eq!(budget.active_tokens(), 800);
}
}