use std::collections::HashMap;
use std::future::Future;
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::task::{Context, Poll};
use std::time::{Duration, SystemTime};
use arc_swap::ArcSwap;
use dashmap::DashMap;
use tower::{Layer, Service};
use super::cost::observe_stream_usage;
use super::types::{LlmRequest, LlmRequestKind, LlmResponse};
use crate::client::BoxFuture;
use crate::cost;
use crate::error::{LiterLlmError, Result};
use crate::types::Usage;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum BudgetDimension {
Global,
Model(String),
Tenant(String),
User(String),
ApiKey(String),
}
#[derive(Debug, Clone)]
pub enum BudgetVerdict {
Allow,
Reject {
reason: String,
dimension: BudgetDimension,
},
}
pub struct CostRecordContext<'a> {
pub model: &'a str,
pub provider: &'a str,
pub tenant_id: Option<&'a str>,
pub user_id: Option<&'a str>,
pub api_key_id: Option<&'a str>,
pub cost_usd: f64,
pub tokens_in: u64,
pub tokens_out: u64,
pub timestamp: SystemTime,
}
pub struct CostCheckContext<'a> {
pub model: &'a str,
pub provider: &'a str,
pub tenant_id: Option<&'a str>,
pub user_id: Option<&'a str>,
pub api_key_id: Option<&'a str>,
pub timestamp: SystemTime,
}
#[derive(Debug, Clone, Default)]
pub struct BudgetSnapshot {
pub global_spend_usd: f64,
pub per_model: HashMap<String, f64>,
pub per_tenant: HashMap<String, f64>,
pub per_user: HashMap<String, f64>,
pub per_api_key: HashMap<String, f64>,
pub limit_global: Option<f64>,
pub limits_per_user: HashMap<String, f64>,
pub limits_per_api_key: HashMap<String, f64>,
pub limits_per_tenant: HashMap<String, f64>,
}
pub trait BudgetLedger: Send + Sync + 'static {
fn record<'a>(&'a self, ctx: &'a CostRecordContext<'a>) -> Pin<Box<dyn Future<Output = ()> + Send + 'a>>;
fn check<'a>(&'a self, ctx: &'a CostCheckContext<'a>) -> Pin<Box<dyn Future<Output = BudgetVerdict> + Send + 'a>>;
fn snapshot(&self) -> BudgetSnapshot;
}
#[derive(Debug)]
struct WindowEntry {
spend_mc: AtomicU64,
window_start_secs: AtomicU64,
window_secs: u64,
}
impl WindowEntry {
fn new(window: Duration) -> Self {
let now_secs = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
Self {
spend_mc: AtomicU64::new(0),
window_start_secs: AtomicU64::new(now_secs),
window_secs: window.as_secs(),
}
}
fn spend_usd(&self, now: SystemTime) -> f64 {
let now_secs = now.duration_since(SystemTime::UNIX_EPOCH).unwrap_or_default().as_secs();
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);
}
}
microcents_to_usd(self.spend_mc.load(Ordering::Acquire))
}
fn add(&self, usd: f64, now: SystemTime) {
let _ = self.spend_usd(now);
self.spend_mc.fetch_add(usd_to_microcents(usd), Ordering::AcqRel);
}
}
#[derive(Debug, Clone, Default)]
pub struct DimensionLimits {
pub global: Option<f64>,
pub per_model: HashMap<String, f64>,
pub per_tenant: HashMap<String, f64>,
pub per_user: HashMap<String, f64>,
pub per_api_key: HashMap<String, f64>,
}
#[derive(Debug)]
pub struct InMemoryBudgetLedger {
limits: ArcSwap<DimensionLimits>,
window: Duration,
global: Arc<WindowEntry>,
per_model: Arc<DashMap<String, WindowEntry>>,
per_tenant: Arc<DashMap<String, WindowEntry>>,
per_user: Arc<DashMap<String, WindowEntry>>,
per_api_key: Arc<DashMap<String, WindowEntry>>,
max_principal_entries: usize,
}
impl InMemoryBudgetLedger {
pub const DEFAULT_MAX_PRINCIPAL_ENTRIES: usize = 100_000;
#[must_use]
pub fn new(limits: DimensionLimits, window: Duration) -> Self {
Self {
global: Arc::new(WindowEntry::new(window)),
per_model: Arc::new(DashMap::new()),
per_tenant: Arc::new(DashMap::new()),
per_user: Arc::new(DashMap::new()),
per_api_key: Arc::new(DashMap::new()),
limits: ArcSwap::from_pointee(limits),
window,
max_principal_entries: Self::DEFAULT_MAX_PRINCIPAL_ENTRIES,
}
}
#[must_use]
pub fn with_max_principal_entries(mut self, max_principal_entries: usize) -> Self {
self.max_principal_entries = max_principal_entries.max(1);
self
}
pub fn update_limits(&self, limits: DimensionLimits) {
self.limits.store(Arc::new(limits));
}
#[must_use]
pub fn from_config(config: &BudgetConfig) -> Self {
let limits = DimensionLimits {
global: config.global_limit,
per_model: config.model_limits.clone(),
..Default::default()
};
Self::new(limits, Duration::from_secs(30 * 24 * 3600))
}
pub fn export_csv(&self, mut writer: impl io::Write) -> io::Result<()> {
let snap = self.snapshot();
writeln!(writer, "dimension,spend_usd")?;
writeln!(writer, "global,{}", snap.global_spend_usd)?;
for (model, spend) in &snap.per_model {
writeln!(writer, "model:{model},{spend}")?;
}
for (tenant, spend) in &snap.per_tenant {
writeln!(writer, "tenant:{tenant},{spend}")?;
}
for (user, spend) in &snap.per_user {
writeln!(writer, "user:{user},{spend}")?;
}
for (key, spend) in &snap.per_api_key {
writeln!(writer, "api_key:{key},{spend}")?;
}
Ok(())
}
pub fn reset(&self) {
let now = SystemTime::now();
let zero_secs = SystemTime::UNIX_EPOCH
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
self.global.spend_mc.store(0, Ordering::Relaxed);
self.global.window_start_secs.store(zero_secs, Ordering::Relaxed);
let _ = self.global.spend_usd(now);
self.per_model.clear();
self.per_tenant.clear();
self.per_user.clear();
self.per_api_key.clear();
}
fn entry_spend(map: &DashMap<String, WindowEntry>, key: &str, now: SystemTime) -> f64 {
map.get(key).map(|e| e.spend_usd(now)).unwrap_or(0.0)
}
fn entry_add(map: &DashMap<String, WindowEntry>, key: &str, usd: f64, window: Duration, now: SystemTime) {
map.entry(key.to_owned())
.or_insert_with(|| WindowEntry::new(window))
.add(usd, now);
}
fn entry_add_capped(
map: &DashMap<String, WindowEntry>,
dimension_name: &'static str,
key: &str,
usd: f64,
window: Duration,
now: SystemTime,
max_entries: usize,
) {
if !map.contains_key(key) && map.len() >= max_entries {
Self::reclaim_zero_spend_entries(map, dimension_name, now);
}
if !map.contains_key(key) && map.len() >= max_entries {
tracing::warn!(
dimension = dimension_name,
cap = max_entries,
"budget ledger at capacity; declining to track new principal, spend for this call was not recorded"
);
return;
}
map.entry(key.to_owned())
.or_insert_with(|| WindowEntry::new(window))
.add(usd, now);
}
fn reclaim_zero_spend_entries(map: &DashMap<String, WindowEntry>, dimension_name: &'static str, now: SystemTime) {
let mut evicted_live_spend_usd = 0.0_f64;
map.retain(|_, entry| {
let spend = entry.spend_usd(now);
if spend > 0.0 {
return true;
}
evicted_live_spend_usd += spend;
false
});
if evicted_live_spend_usd > 0.0 {
tracing::warn!(
dimension = dimension_name,
evicted_spend_usd = evicted_live_spend_usd,
"budget ledger eviction reclaimed entries carrying non-zero spend; budget enforcement was weakened"
);
}
}
fn check_limit(spend: f64, limit: f64, dimension: BudgetDimension, key: &str) -> Option<BudgetVerdict> {
if spend >= limit {
Some(BudgetVerdict::Reject {
reason: format!("{key} budget exceeded: spent ${spend:.6}, limit ${limit:.6}"),
dimension,
})
} else {
None
}
}
}
impl BudgetLedger for InMemoryBudgetLedger {
fn record<'a>(&'a self, ctx: &'a CostRecordContext<'a>) -> Pin<Box<dyn Future<Output = ()> + Send + 'a>> {
Box::pin(async move {
let now = ctx.timestamp;
self.global.add(ctx.cost_usd, now);
Self::entry_add(&self.per_model, ctx.model, ctx.cost_usd, self.window, now);
if let Some(tenant) = ctx.tenant_id {
Self::entry_add(&self.per_tenant, tenant, ctx.cost_usd, self.window, now);
}
if let Some(user) = ctx.user_id {
Self::entry_add_capped(
&self.per_user,
"user",
user,
ctx.cost_usd,
self.window,
now,
self.max_principal_entries,
);
}
if let Some(key) = ctx.api_key_id {
Self::entry_add_capped(
&self.per_api_key,
"api_key",
key,
ctx.cost_usd,
self.window,
now,
self.max_principal_entries,
);
}
#[cfg(feature = "otel")]
{
use super::metrics;
metrics::record_budget_spend(
ctx.model,
ctx.provider,
ctx.tenant_id,
ctx.user_id,
ctx.api_key_id,
ctx.cost_usd,
);
}
})
}
fn check<'a>(&'a self, ctx: &'a CostCheckContext<'a>) -> Pin<Box<dyn Future<Output = BudgetVerdict> + Send + 'a>> {
Box::pin(async move {
let now = ctx.timestamp;
let limits = self.limits.load_full();
if let Some(limit) = limits.global {
let spend = self.global.spend_usd(now);
if let Some(v) = Self::check_limit(spend, limit, BudgetDimension::Global, "global") {
return v;
}
}
if let Some(&limit) = limits.per_model.get(ctx.model) {
let spend = Self::entry_spend(&self.per_model, ctx.model, now);
if let Some(v) = Self::check_limit(
spend,
limit,
BudgetDimension::Model(ctx.model.to_owned()),
&format!("model:{}", ctx.model),
) {
return v;
}
}
if let Some(tenant) = ctx.tenant_id
&& let Some(&limit) = limits.per_tenant.get(tenant)
{
let spend = Self::entry_spend(&self.per_tenant, tenant, now);
if let Some(v) = Self::check_limit(
spend,
limit,
BudgetDimension::Tenant(tenant.to_owned()),
&format!("tenant:{tenant}"),
) {
return v;
}
}
if let Some(user) = ctx.user_id
&& let Some(&limit) = limits.per_user.get(user)
{
let spend = Self::entry_spend(&self.per_user, user, now);
if let Some(v) = Self::check_limit(
spend,
limit,
BudgetDimension::User(user.to_owned()),
&format!("user:{user}"),
) {
return v;
}
}
if let Some(key) = ctx.api_key_id
&& let Some(&limit) = limits.per_api_key.get(key)
{
let spend = Self::entry_spend(&self.per_api_key, key, now);
if let Some(v) = Self::check_limit(
spend,
limit,
BudgetDimension::ApiKey(key.to_owned()),
&format!("api_key:{key}"),
) {
return v;
}
}
BudgetVerdict::Allow
})
}
fn snapshot(&self) -> BudgetSnapshot {
let now = SystemTime::now();
let limits = self.limits.load();
let global_spend_usd = self.global.spend_usd(now);
let per_model = self
.per_model
.iter()
.map(|e| (e.key().clone(), e.value().spend_usd(now)))
.collect();
let per_tenant = self
.per_tenant
.iter()
.map(|e| (e.key().clone(), e.value().spend_usd(now)))
.collect();
let per_user = self
.per_user
.iter()
.map(|e| (e.key().clone(), e.value().spend_usd(now)))
.collect();
let per_api_key = self
.per_api_key
.iter()
.map(|e| (e.key().clone(), e.value().spend_usd(now)))
.collect();
BudgetSnapshot {
global_spend_usd,
per_model,
per_tenant,
per_user,
per_api_key,
limit_global: limits.global,
limits_per_user: limits.per_user.clone(),
limits_per_api_key: limits.per_api_key.clone(),
limits_per_tenant: limits.per_tenant.clone(),
}
}
}
#[must_use]
pub fn should_hedge<L: BudgetLedger>(
ledger: &L,
ctx: &CostCheckContext<'_>,
estimated_cost_usd: f64,
safety_margin_pct: f64,
) -> bool {
let snap = ledger.snapshot();
let hedge_cost = 2.0 * estimated_cost_usd;
let margin = safety_margin_pct.clamp(0.0, 0.999);
let has_headroom = |spend: f64, limit: f64| -> bool {
let effective_limit = limit * (1.0 - margin);
spend + hedge_cost < effective_limit
};
if let Some(global_limit) = snap.limit_global
&& !has_headroom(snap.global_spend_usd, global_limit)
{
return false;
}
if let Some(user) = ctx.user_id
&& let Some(&user_limit) = snap.limits_per_user.get(user)
{
let user_spend = snap.per_user.get(user).copied().unwrap_or(0.0);
if !has_headroom(user_spend, user_limit) {
return false;
}
}
if let Some(key) = ctx.api_key_id
&& let Some(&key_limit) = snap.limits_per_api_key.get(key)
{
let key_spend = snap.per_api_key.get(key).copied().unwrap_or(0.0);
if !has_headroom(key_spend, key_limit) {
return false;
}
}
if let Some(tenant) = ctx.tenant_id
&& let Some(&tenant_limit) = snap.limits_per_tenant.get(tenant)
{
let tenant_spend = snap.per_tenant.get(tenant).copied().unwrap_or(0.0);
if !has_headroom(tenant_spend, tenant_limit) {
return false;
}
}
true
}
pub(crate) fn provider_of(model: &str) -> &str {
model.split_once('/').map_or("", |(prefix, _)| prefix)
}
pub(crate) fn user_id_of(req: &LlmRequest) -> Option<&str> {
match &req.kind {
LlmRequestKind::Chat(r) | LlmRequestKind::ChatStream(r) => r.user.as_deref(),
LlmRequestKind::Embed(r) => r.user.as_deref(),
_ => None,
}
}
#[cfg_attr(alef, alef(skip))]
pub struct BudgetLedgerLayer<L: BudgetLedger> {
ledger: Arc<L>,
enforcement: Enforcement,
}
impl<L: BudgetLedger> BudgetLedgerLayer<L> {
#[cfg_attr(alef, alef(skip))]
#[must_use]
pub fn new(ledger: Arc<L>, enforcement: Enforcement) -> Self {
Self { ledger, enforcement }
}
}
impl<L: BudgetLedger, S> Layer<S> for BudgetLedgerLayer<L> {
type Service = BudgetLedgerService<L, S>;
fn layer(&self, inner: S) -> Self::Service {
BudgetLedgerService {
inner,
ledger: Arc::clone(&self.ledger),
enforcement: self.enforcement,
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct BudgetLedgerService<L: BudgetLedger, S> {
inner: S,
ledger: Arc<L>,
enforcement: Enforcement,
}
impl<L: BudgetLedger, S: Clone> Clone for BudgetLedgerService<L, S> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
ledger: Arc::clone(&self.ledger),
enforcement: self.enforcement,
}
}
}
impl<L, S> Service<LlmRequest> for BudgetLedgerService<L, S>
where
L: BudgetLedger,
S: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Clone + 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 provider = provider_of(&model).to_owned();
let tenant_id = req.tenant_id().map(|t| t.as_ref().to_owned());
let user_id = user_id_of(&req).map(str::to_owned);
let ledger = Arc::clone(&self.ledger);
let enforcement = self.enforcement;
let standby = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, standby);
Box::pin(async move {
let check_ctx = CostCheckContext {
model: &model,
provider: &provider,
tenant_id: tenant_id.as_deref(),
user_id: user_id.as_deref(),
api_key_id: None,
timestamp: SystemTime::now(),
};
if enforcement == Enforcement::Hard
&& let BudgetVerdict::Reject { reason, dimension } = ledger.check(&check_ctx).await
{
let model_field = match &dimension {
BudgetDimension::Model(m) => Some(m.clone()),
_ => None,
};
return Err(LiterLlmError::BudgetExceeded {
message: reason,
model: model_field,
});
}
let resp = inner.call(req).await?;
match resp {
LlmResponse::ChatStream(stream) => {
let ledger_for_completion = Arc::clone(&ledger);
let model_c = model.clone();
let provider_c = provider.clone();
let tenant_c = tenant_id.clone();
let user_c = user_id.clone();
let wrapped = observe_stream_usage(stream, move |usage| {
let Some(usage) = usage else { return };
let Some(usd) = cost::completion_cost(&model_c, usage.prompt_tokens, usage.completion_tokens)
else {
return;
};
let record = async move {
let ctx = CostRecordContext {
model: &model_c,
provider: &provider_c,
tenant_id: tenant_c.as_deref(),
user_id: user_c.as_deref(),
api_key_id: None,
cost_usd: usd,
tokens_in: usage.prompt_tokens,
tokens_out: usage.completion_tokens,
timestamp: SystemTime::now(),
};
ledger_for_completion.record(&ctx).await;
};
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(record);
} else {
tracing::warn!(
"budget ledger: no Tokio runtime available to record streamed usage; spend was not recorded"
);
}
});
Ok(LlmResponse::ChatStream(wrapped))
}
other => {
if let Some(usage) = other.usage()
&& let Some(usd) = cost::completion_cost(&model, usage.prompt_tokens, usage.completion_tokens)
{
let ctx = CostRecordContext {
model: &model,
provider: &provider,
tenant_id: tenant_id.as_deref(),
user_id: user_id.as_deref(),
api_key_id: None,
cost_usd: usd,
tokens_in: usage.prompt_tokens,
tokens_out: usage.completion_tokens,
timestamp: SystemTime::now(),
};
ledger.record(&ctx).await;
}
Ok(other)
}
}
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum Enforcement {
Hard,
Soft,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct BudgetConfig {
pub global_limit: Option<f64>,
pub model_limits: HashMap<String, f64>,
pub enforcement: Enforcement,
}
impl Default for BudgetConfig {
fn default() -> Self {
Self {
global_limit: None,
model_limits: HashMap::new(),
enforcement: Enforcement::Hard,
}
}
}
#[derive(Debug)]
pub struct BudgetState {
global_spend: AtomicU64,
model_spend: DashMap<String, AtomicU64>,
}
impl BudgetState {
#[must_use]
pub fn new() -> Self {
Self {
global_spend: AtomicU64::new(0),
model_spend: DashMap::new(),
}
}
#[must_use]
pub fn global_spend(&self) -> f64 {
microcents_to_usd(self.global_spend.load(Ordering::Relaxed))
}
#[must_use]
pub fn model_spend(&self, model: &str) -> f64 {
self.model_spend
.get(model)
.map(|v| microcents_to_usd(v.load(Ordering::Relaxed)))
.unwrap_or(0.0)
}
pub fn reset(&self) {
self.global_spend.store(0, Ordering::Relaxed);
self.model_spend.clear();
}
fn record(&self, model: &str, usd: f64) {
let mc = usd_to_microcents(usd);
self.global_spend.fetch_add(mc, Ordering::Relaxed);
self.model_spend
.entry(model.to_owned())
.or_insert_with(|| AtomicU64::new(0))
.fetch_add(mc, Ordering::Relaxed);
}
}
#[cfg_attr(alef, alef(skip))]
impl Default for BudgetState {
fn default() -> Self {
Self::new()
}
}
fn usd_to_microcents(usd: f64) -> u64 {
if usd <= 0.0 {
return 0;
}
(usd * 1_000_000.0).round() as u64
}
fn microcents_to_usd(mc: u64) -> f64 {
mc as f64 / 1_000_000.0
}
#[cfg_attr(alef, alef(skip))]
pub struct BudgetLayer {
config: BudgetConfig,
state: Arc<BudgetState>,
}
#[cfg_attr(alef, alef(skip))]
impl BudgetLayer {
#[must_use]
pub fn new(config: BudgetConfig, state: Arc<BudgetState>) -> Self {
Self { config, state }
}
}
impl<S> Layer<S> for BudgetLayer {
type Service = BudgetService<S>;
fn layer(&self, inner: S) -> Self::Service {
BudgetService {
inner,
config: self.config.clone(),
state: Arc::clone(&self.state),
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct BudgetService<S> {
inner: S,
config: BudgetConfig,
state: Arc<BudgetState>,
}
impl<S: Clone> Clone for BudgetService<S> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
config: self.config.clone(),
state: Arc::clone(&self.state),
}
}
}
impl<S> Service<LlmRequest> for BudgetService<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 config.enforcement == Enforcement::Hard
&& let Some(err) = check_budget(&config, &state, &model)
{
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 config_for_completion = config.clone();
let wrapped = observe_stream_usage(stream, move |usage| {
record_usage(
&config_for_completion,
&state_for_completion,
&model_for_completion,
usage.as_ref(),
);
});
Ok(LlmResponse::ChatStream(wrapped))
}
other => {
record_usage(&config, &state, &model, other.usage());
Ok(other)
}
}
})
}
}
fn record_usage(config: &BudgetConfig, state: &BudgetState, model: &str, usage: Option<&Usage>) {
let Some(usage) = usage else { return };
let Some(usd) = cost::completion_cost(model, usage.prompt_tokens, usage.completion_tokens) else {
return;
};
state.record(model, usd);
if config.enforcement == Enforcement::Soft {
emit_soft_warnings(config, state, model);
}
}
fn check_budget(config: &BudgetConfig, state: &BudgetState, model: &str) -> Option<LiterLlmError> {
if let Some(limit) = config.global_limit
&& state.global_spend() >= limit
{
return Some(LiterLlmError::BudgetExceeded {
message: format!(
"global budget exceeded: spent ${:.6}, limit ${:.6}",
state.global_spend(),
limit,
),
model: None,
});
}
if let Some(&limit) = config.model_limits.get(model)
&& state.model_spend(model) >= limit
{
return Some(LiterLlmError::BudgetExceeded {
message: format!(
"model {model} budget exceeded: spent ${:.6}, limit ${:.6}",
state.model_spend(model),
limit,
),
model: Some(model.to_owned()),
});
}
None
}
fn emit_soft_warnings(config: &BudgetConfig, state: &BudgetState, model: &str) {
if let Some(limit) = config.global_limit
&& state.global_spend() >= limit
{
tracing::warn!(
spend = state.global_spend(),
limit,
"global budget exceeded (soft enforcement)"
);
}
if let Some(&limit) = config.model_limits.get(model)
&& state.model_spend(model) >= limit
{
tracing::warn!(
model,
spend = state.model_spend(model),
limit,
"model budget exceeded (soft enforcement)"
);
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use super::*;
use crate::tower::service::LlmService;
use crate::tower::tests_common::{MockClient, chat_req};
use crate::tower::types::LlmRequest;
fn build_service(config: BudgetConfig, state: Arc<BudgetState>) -> BudgetService<LlmService<MockClient>> {
let layer = BudgetLayer::new(config, state);
let inner = LlmService::new(MockClient::ok());
layer.layer(inner)
}
#[tokio::test]
async fn hard_enforcement_rejects_when_global_limit_exceeded() {
let state = Arc::new(BudgetState::new());
state.global_spend.store(usd_to_microcents(10.0), Ordering::Relaxed);
let config = BudgetConfig {
global_limit: Some(5.0),
enforcement: Enforcement::Hard,
..Default::default()
};
let mut svc = build_service(config, state);
let err = svc
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect_err("should reject over-budget request");
assert!(matches!(err, LiterLlmError::BudgetExceeded { .. }));
}
#[tokio::test]
async fn hard_enforcement_rejects_when_model_limit_exceeded() {
let state = Arc::new(BudgetState::new());
state
.model_spend
.entry("gpt-4".to_owned())
.or_insert_with(|| AtomicU64::new(0))
.store(usd_to_microcents(2.0), Ordering::Relaxed);
let mut limits = HashMap::new();
limits.insert("gpt-4".into(), 1.0);
let config = BudgetConfig {
global_limit: None,
model_limits: limits,
enforcement: Enforcement::Hard,
};
let mut svc = build_service(config, state);
let err = svc
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect_err("should reject over-budget model request");
match &err {
LiterLlmError::BudgetExceeded { model, .. } => {
assert_eq!(model.as_deref(), Some("gpt-4"));
}
other => panic!("expected BudgetExceeded, got {other:?}"),
}
}
#[tokio::test]
async fn hard_enforcement_allows_requests_under_limit() {
let state = Arc::new(BudgetState::new());
let config = BudgetConfig {
global_limit: Some(100.0),
enforcement: Enforcement::Hard,
..Default::default()
};
let mut svc = build_service(config, state);
let resp = svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok(), "request under budget should succeed");
}
#[tokio::test]
async fn soft_enforcement_allows_requests_over_global_limit() {
let state = Arc::new(BudgetState::new());
state.global_spend.store(usd_to_microcents(100.0), Ordering::Relaxed);
let config = BudgetConfig {
global_limit: Some(5.0),
enforcement: Enforcement::Soft,
..Default::default()
};
let mut svc = build_service(config, state);
let resp = svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok(), "soft mode should never reject");
}
#[tokio::test]
async fn soft_enforcement_allows_requests_over_model_limit() {
let state = Arc::new(BudgetState::new());
state
.model_spend
.entry("gpt-4".to_owned())
.or_insert_with(|| AtomicU64::new(0))
.store(usd_to_microcents(10.0), Ordering::Relaxed);
let mut limits = HashMap::new();
limits.insert("gpt-4".into(), 1.0);
let config = BudgetConfig {
global_limit: None,
model_limits: limits,
enforcement: Enforcement::Soft,
};
let mut svc = build_service(config, state);
let resp = svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok(), "soft mode should never reject");
}
#[tokio::test]
async fn accumulates_cost_after_response() {
let state = Arc::new(BudgetState::new());
let config = BudgetConfig {
global_limit: Some(100.0),
enforcement: Enforcement::Hard,
..Default::default()
};
let mut svc = build_service(config, Arc::clone(&state));
svc.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("service call should not fail");
assert!(state.global_spend() > 0.0, "global spend should be recorded");
assert!(state.model_spend("gpt-4") > 0.0, "model spend should be recorded");
}
#[tokio::test]
async fn per_model_limits_are_independent() {
let state = Arc::new(BudgetState::new());
state
.model_spend
.entry("gpt-4".to_owned())
.or_insert_with(|| AtomicU64::new(0))
.store(usd_to_microcents(5.0), Ordering::Relaxed);
let mut limits = HashMap::new();
limits.insert("gpt-4".into(), 1.0);
let config = BudgetConfig {
global_limit: None,
model_limits: limits,
enforcement: Enforcement::Hard,
};
let mut svc = build_service(config, state);
let err = svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(err.is_err(), "gpt-4 should be rejected");
let ok = svc.call(LlmRequest::Chat(chat_req("gpt-3.5-turbo"))).await;
assert!(ok.is_ok(), "gpt-3.5-turbo should not be limited");
}
#[tokio::test]
async fn reset_clears_all_counters() {
let state = Arc::new(BudgetState::new());
state.global_spend.store(usd_to_microcents(50.0), Ordering::Relaxed);
state
.model_spend
.entry("gpt-4".to_owned())
.or_insert_with(|| AtomicU64::new(0))
.store(usd_to_microcents(25.0), Ordering::Relaxed);
assert!(state.global_spend() > 0.0);
assert!(state.model_spend("gpt-4") > 0.0);
state.reset();
assert_eq!(state.global_spend(), 0.0, "global spend should be zero after reset");
assert_eq!(
state.model_spend("gpt-4"),
0.0,
"model spend should be zero after reset"
);
}
#[tokio::test]
async fn reset_allows_previously_blocked_requests() {
let state = Arc::new(BudgetState::new());
state.global_spend.store(usd_to_microcents(10.0), Ordering::Relaxed);
let config = BudgetConfig {
global_limit: Some(5.0),
enforcement: Enforcement::Hard,
..Default::default()
};
let mut svc = build_service(config, Arc::clone(&state));
let err = svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(err.is_err());
state.reset();
let ok = svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(ok.is_ok(), "should succeed after reset");
}
#[tokio::test]
async fn unlimited_config_allows_all_requests() {
let state = Arc::new(BudgetState::new());
let config = BudgetConfig::default();
let mut svc = build_service(config, state);
for _ in 0..20 {
assert!(svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await.is_ok());
}
}
#[tokio::test]
async fn propagates_inner_service_errors() {
let state = Arc::new(BudgetState::new());
let config = BudgetConfig {
global_limit: Some(100.0),
enforcement: Enforcement::Hard,
..Default::default()
};
let layer = BudgetLayer::new(config, state);
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("should propagate inner error");
assert!(matches!(err, LiterLlmError::Timeout));
}
#[tokio::test]
async fn budget_ledger_records_per_key_and_per_user() {
let limits = DimensionLimits::default();
let ledger = InMemoryBudgetLedger::new(limits, Duration::from_secs(3600));
let ctx1 = CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: Some("acme"),
user_id: Some("alice"),
api_key_id: Some("key-1"),
cost_usd: 0.10,
tokens_in: 1000,
tokens_out: 500,
timestamp: SystemTime::now(),
};
ledger.record(&ctx1).await;
let ctx2 = CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: Some("acme"),
user_id: Some("bob"),
api_key_id: Some("key-2"),
cost_usd: 0.20,
tokens_in: 2000,
tokens_out: 1000,
timestamp: SystemTime::now(),
};
ledger.record(&ctx2).await;
let snap = ledger.snapshot();
assert!(
(snap.global_spend_usd - 0.30).abs() < 1e-9,
"global: {}",
snap.global_spend_usd
);
assert!((snap.per_model["gpt-4"] - 0.30).abs() < 1e-9);
assert!((snap.per_tenant["acme"] - 0.30).abs() < 1e-9);
assert!((snap.per_user["alice"] - 0.10).abs() < 1e-9);
assert!((snap.per_user["bob"] - 0.20).abs() < 1e-9);
assert!((snap.per_api_key["key-1"] - 0.10).abs() < 1e-9);
assert!((snap.per_api_key["key-2"] - 0.20).abs() < 1e-9);
}
#[tokio::test]
async fn budget_ledger_per_user_map_is_bounded_by_max_principal_entries() {
const CAP: usize = 8;
let ledger = InMemoryBudgetLedger::new(DimensionLimits::default(), Duration::from_secs(3600))
.with_max_principal_entries(CAP);
for i in 0..CAP * 10 {
let user = format!("user-{i}");
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some(&user),
api_key_id: None,
cost_usd: 0.01,
tokens_in: 10,
tokens_out: 5,
timestamp: SystemTime::now(),
})
.await;
}
let len = ledger.per_user.len();
assert!(
len <= CAP,
"per_user must stay within its cap; held {len} entries with a cap of {CAP}"
);
}
#[tokio::test]
async fn budget_ledger_per_api_key_map_is_bounded_by_max_principal_entries() {
const CAP: usize = 8;
let ledger = InMemoryBudgetLedger::new(DimensionLimits::default(), Duration::from_secs(3600))
.with_max_principal_entries(CAP);
for i in 0..CAP * 10 {
let key = format!("key-{i}");
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: None,
api_key_id: Some(&key),
cost_usd: 0.01,
tokens_in: 10,
tokens_out: 5,
timestamp: SystemTime::now(),
})
.await;
}
let len = ledger.per_api_key.len();
assert!(
len <= CAP,
"per_api_key must stay within its cap; held {len} entries with a cap of {CAP}"
);
}
#[tokio::test]
async fn budget_ledger_cap_reclaims_zero_spend_entries_before_declining_new_principals() {
const CAP: usize = 4;
let ledger = InMemoryBudgetLedger::new(DimensionLimits::default(), Duration::from_secs(3600))
.with_max_principal_entries(CAP);
for i in 0..CAP {
let user = format!("zero-spend-{i}");
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some(&user),
api_key_id: None,
cost_usd: 0.0,
tokens_in: 0,
tokens_out: 0,
timestamp: SystemTime::now(),
})
.await;
}
assert_eq!(ledger.per_user.len(), CAP, "setup should fill the ledger to its cap");
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("real-spender"),
api_key_id: None,
cost_usd: 5.0,
tokens_in: 100,
tokens_out: 50,
timestamp: SystemTime::now(),
})
.await;
assert_eq!(
ledger.per_user.len(),
1,
"the zero-spend entries must be reclaimed, leaving only the new spender"
);
assert!((ledger.snapshot().per_user["real-spender"] - 5.0).abs() < 1e-9);
}
#[tokio::test]
async fn budget_ledger_cap_never_resets_an_existing_principals_spend() {
const CAP: usize = 4;
let ledger = InMemoryBudgetLedger::new(DimensionLimits::default(), Duration::from_secs(3600))
.with_max_principal_entries(CAP);
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
cost_usd: 42.0,
tokens_in: 100,
tokens_out: 50,
timestamp: SystemTime::now(),
})
.await;
assert!((ledger.snapshot().per_user["alice"] - 42.0).abs() < 1e-9);
for i in 0..CAP * 10 {
let user = format!("flood-{i}");
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some(&user),
api_key_id: None,
cost_usd: 1.0,
tokens_in: 10,
tokens_out: 5,
timestamp: SystemTime::now(),
})
.await;
}
let spend_after_flood = ledger.snapshot().per_user.get("alice").copied();
assert!(
matches!(spend_after_flood, Some(v) if (v - 42.0).abs() < 1e-9),
"alice's already-recorded spend must survive cap pressure from other principals, got {spend_after_flood:?}"
);
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
cost_usd: 8.0,
tokens_in: 10,
tokens_out: 5,
timestamp: SystemTime::now(),
})
.await;
assert!(
(ledger.snapshot().per_user["alice"] - 50.0).abs() < 1e-9,
"an already-tracked principal must keep accumulating spend once the ledger is at capacity"
);
}
#[tokio::test]
async fn update_limits_preserves_accumulated_spend() {
let mut limits = DimensionLimits::default();
limits.per_user.insert("alice".to_owned(), 100.0);
let ledger = InMemoryBudgetLedger::new(limits, Duration::from_secs(3600));
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
cost_usd: 7.50,
tokens_in: 100,
tokens_out: 50,
timestamp: SystemTime::now(),
})
.await;
assert!((ledger.snapshot().per_user["alice"] - 7.50).abs() < 1e-9);
let mut new_limits = DimensionLimits::default();
new_limits.per_user.insert("alice".to_owned(), 200.0);
ledger.update_limits(new_limits);
let snap = ledger.snapshot();
assert!(
(snap.per_user["alice"] - 7.50).abs() < 1e-9,
"spend must survive update_limits, got {}",
snap.per_user["alice"]
);
assert_eq!(snap.limits_per_user.get("alice"), Some(&200.0));
let verdict = ledger
.check(&CostCheckContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
timestamp: SystemTime::now(),
})
.await;
assert!(
matches!(verdict, BudgetVerdict::Allow),
"spend $7.50 is well under the new $200 limit, expected Allow, got {verdict:?}"
);
}
#[tokio::test]
async fn update_limits_lowering_below_spend_rejects_immediately() {
let mut limits = DimensionLimits::default();
limits.per_user.insert("alice".to_owned(), 100.0);
let ledger = InMemoryBudgetLedger::new(limits, Duration::from_secs(3600));
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
cost_usd: 5.0,
tokens_in: 100,
tokens_out: 50,
timestamp: SystemTime::now(),
})
.await;
let mut lowered = DimensionLimits::default();
lowered.per_user.insert("alice".to_owned(), 1.0);
ledger.update_limits(lowered);
let verdict = ledger
.check(&CostCheckContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
timestamp: SystemTime::now(),
})
.await;
match verdict {
BudgetVerdict::Reject { dimension, .. } => {
assert!(matches!(dimension, BudgetDimension::User(ref u) if u == "alice"));
}
BudgetVerdict::Allow => panic!("lowering the limit below existing spend must reject, got Allow"),
}
}
#[tokio::test]
async fn update_limits_retains_spend_for_removed_tenant() {
let mut limits = DimensionLimits::default();
limits.per_tenant.insert("acme".to_owned(), 10.0);
let ledger = InMemoryBudgetLedger::new(limits, Duration::from_secs(3600));
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: Some("acme"),
user_id: None,
api_key_id: None,
cost_usd: 3.0,
tokens_in: 100,
tokens_out: 50,
timestamp: SystemTime::now(),
})
.await;
ledger.update_limits(DimensionLimits::default());
let verdict = ledger
.check(&CostCheckContext {
model: "gpt-4",
provider: "openai",
tenant_id: Some("acme"),
user_id: None,
api_key_id: None,
timestamp: SystemTime::now(),
})
.await;
assert!(
matches!(verdict, BudgetVerdict::Allow),
"no configured limit means no enforcement, expected Allow, got {verdict:?}"
);
let snap = ledger.snapshot();
assert!(
(snap.per_tenant["acme"] - 3.0).abs() < 1e-9,
"spend history for a removed tenant must be retained, got {:?}",
snap.per_tenant.get("acme")
);
}
#[tokio::test]
async fn budget_ledger_rejects_when_user_limit_exceeded() {
let mut limits = DimensionLimits::default();
limits.per_user.insert("alice".to_owned(), 0.05);
let ledger = InMemoryBudgetLedger::new(limits, Duration::from_secs(3600));
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
cost_usd: 0.10,
tokens_in: 100,
tokens_out: 50,
timestamp: SystemTime::now(),
})
.await;
let verdict = ledger
.check(&CostCheckContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
timestamp: SystemTime::now(),
})
.await;
match verdict {
BudgetVerdict::Reject { dimension, .. } => {
assert!(
matches!(dimension, BudgetDimension::User(ref u) if u == "alice"),
"expected User(alice) dimension, got {dimension:?}"
);
}
BudgetVerdict::Allow => panic!("expected Reject, got Allow"),
}
}
#[tokio::test]
async fn budget_ledger_resets_at_window_boundary() {
let limits = DimensionLimits {
global: Some(100.0),
..Default::default()
};
let window = Duration::from_secs(1);
let ledger = InMemoryBudgetLedger::new(limits, window);
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: None,
api_key_id: None,
cost_usd: 50.0,
tokens_in: 1_000_000,
tokens_out: 0,
timestamp: SystemTime::now(),
})
.await;
assert!(ledger.snapshot().global_spend_usd > 0.0);
let future = SystemTime::now() + Duration::from_secs(2);
let spend_after_window = ledger.global.spend_usd(future);
assert_eq!(spend_after_window, 0.0, "spend should reset to 0 after window boundary");
}
#[tokio::test]
async fn budget_snapshot_csv_export_round_trips() {
let ledger = InMemoryBudgetLedger::new(DimensionLimits::default(), Duration::from_secs(3600));
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: Some("tenant-x"),
user_id: Some("user-y"),
api_key_id: Some("key-z"),
cost_usd: 1.23,
tokens_in: 100,
tokens_out: 50,
timestamp: SystemTime::now(),
})
.await;
let mut csv_bytes: Vec<u8> = Vec::new();
ledger.export_csv(&mut csv_bytes).expect("CSV export must not fail");
let csv = String::from_utf8(csv_bytes).expect("CSV must be valid UTF-8");
assert!(csv.starts_with("dimension,spend_usd\n"), "missing header: {csv}");
let mut found_global = false;
let mut found_model = false;
let mut found_tenant = false;
let mut found_user = false;
let mut found_key = false;
for line in csv.lines().skip(1) {
let parts: Vec<&str> = line.splitn(2, ',').collect();
assert_eq!(parts.len(), 2, "malformed CSV line: {line}");
let dimension = parts[0];
let spend: f64 = parts[1].parse().expect("spend must be a float");
match dimension {
"global" => {
assert!((spend - 1.23).abs() < 1e-6, "global spend mismatch: {spend}");
found_global = true;
}
"model:gpt-4" => {
assert!((spend - 1.23).abs() < 1e-6);
found_model = true;
}
"tenant:tenant-x" => {
assert!((spend - 1.23).abs() < 1e-6);
found_tenant = true;
}
"user:user-y" => {
assert!((spend - 1.23).abs() < 1e-6);
found_user = true;
}
"api_key:key-z" => {
assert!((spend - 1.23).abs() < 1e-6);
found_key = true;
}
_ => {}
}
}
assert!(found_global, "global row missing from CSV");
assert!(found_model, "model row missing from CSV");
assert!(found_tenant, "tenant row missing from CSV");
assert!(found_user, "user row missing from CSV");
assert!(found_key, "api_key row missing from CSV");
}
#[test]
fn window_rollover_under_concurrent_threads_does_not_undercount() {
use std::sync::Barrier;
use std::thread;
let entry = Arc::new(WindowEntry::new(Duration::from_secs(1)));
let future_now = SystemTime::now() + Duration::from_secs(2);
let barrier = Arc::new(Barrier::new(100));
let mut handles = Vec::with_capacity(100);
for _ in 0..100 {
let entry_clone = Arc::clone(&entry);
let barrier_clone = Arc::clone(&barrier);
handles.push(thread::spawn(move || {
barrier_clone.wait();
entry_clone.add(0.10, future_now);
}));
}
for h in handles {
h.join().expect("thread must not panic");
}
let total = microcents_to_usd(entry.spend_mc.load(Ordering::Acquire));
assert!(
(total - 10.0_f64).abs() < 1e-4,
"expected $10.00 total, got ${total:.6} — window rollover race caused under-counting"
);
}
#[test]
fn budget_window_rollover_no_torn_read() {
use std::sync::Barrier;
use std::thread;
let entry = Arc::new(WindowEntry::new(Duration::from_secs(1)));
let future_now = SystemTime::now() + Duration::from_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 e = Arc::clone(&entry);
let b = Arc::clone(&barrier);
handles.push(thread::spawn(move || {
b.wait();
e.add(0.10, future_now);
}));
}
for h in handles {
h.join().expect("writer must not panic");
}
let total = microcents_to_usd(entry.spend_mc.load(Ordering::Acquire));
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 should_hedge_respects_user_budget() {
let mut limits = DimensionLimits::default();
limits.per_user.insert("alice".to_owned(), 10.0);
let ledger = InMemoryBudgetLedger::new(limits, Duration::from_secs(3600));
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
cost_usd: 9.50,
tokens_in: 100,
tokens_out: 50,
timestamp: SystemTime::now(),
})
.await;
let ctx = CostCheckContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
timestamp: SystemTime::now(),
};
let result = should_hedge(&ledger, &ctx, 0.50, 0.10);
assert!(
!result,
"hedging should be suppressed when user spend + 2×cost would exceed 90% of budget"
);
}
#[tokio::test]
async fn should_hedge_allows_when_far_below_budget() {
let mut limits = DimensionLimits::default();
limits.per_user.insert("alice".to_owned(), 10.0);
let ledger = InMemoryBudgetLedger::new(limits, Duration::from_secs(3600));
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
cost_usd: 1.00,
tokens_in: 100,
tokens_out: 50,
timestamp: SystemTime::now(),
})
.await;
let ctx = CostCheckContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
timestamp: SystemTime::now(),
};
let result = should_hedge(&ledger, &ctx, 0.50, 0.10);
assert!(
result,
"hedging should be allowed when user spend + 2×cost is well below 90% of budget"
);
}
#[derive(Clone)]
struct StreamingUsageService {
prompt_tokens: u64,
completion_tokens: u64,
}
fn usage_chunk(model: &str, usage: Option<Usage>) -> crate::types::ChatCompletionChunk {
crate::types::ChatCompletionChunk {
id: "chunk".into(),
object: "chat.completion.chunk".into(),
created: 0,
model: model.into(),
choices: vec![],
usage,
system_fingerprint: None,
service_tier: None,
}
}
struct ChunkStream(std::collections::VecDeque<crate::types::ChatCompletionChunk>);
impl futures_core::Stream for ChunkStream {
type Item = Result<crate::types::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))
}
}
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 {
let model = req.model().unwrap_or("gpt-4").to_owned();
let usage = Usage {
prompt_tokens: self.prompt_tokens,
completion_tokens: self.completion_tokens,
total_tokens: self.prompt_tokens + self.completion_tokens,
prompt_tokens_details: None,
};
Box::pin(async move {
let chunks =
std::collections::VecDeque::from([usage_chunk(&model, None), usage_chunk(&model, Some(usage))]);
let stream: crate::client::BoxStream<'static, Result<crate::types::ChatCompletionChunk>> =
Box::pin(ChunkStream(chunks));
Ok(LlmResponse::ChatStream(stream))
})
}
}
#[tokio::test]
async fn budget_service_records_cost_for_streamed_response() {
use futures_util::StreamExt as _;
let state = Arc::new(BudgetState::new());
let config = BudgetConfig {
global_limit: Some(1000.0),
enforcement: Enforcement::Hard,
..Default::default()
};
let layer = BudgetLayer::new(config, Arc::clone(&state));
let mut svc = layer.layer(StreamingUsageService {
prompt_tokens: 1000,
completion_tokens: 500,
});
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() {}
assert!(
state.global_spend() > 0.0,
"cost of a streamed response must be recorded in global spend"
);
assert!(
state.model_spend("gpt-4") > 0.0,
"cost of a streamed response must be recorded per-model"
);
}
#[tokio::test]
async fn abandoned_stream_still_records_spend() {
use futures_util::StreamExt as _;
let state = Arc::new(BudgetState::new());
let config = BudgetConfig {
global_limit: Some(1000.0),
enforcement: Enforcement::Hard,
..Default::default()
};
let layer = BudgetLayer::new(config, Arc::clone(&state));
let mut svc = layer.layer(StreamingUsageService {
prompt_tokens: 1000,
completion_tokens: 500,
});
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");
};
let _ = stream.next().await;
drop(stream);
assert!(
state.global_spend() > 0.0,
"spend must be recorded even though the stream was dropped before it ended"
);
assert!(
state.model_spend("gpt-4") > 0.0,
"per-model spend must be recorded for an abandoned stream"
);
}
#[tokio::test]
async fn stream_dropped_without_a_single_poll_records_spend() {
let state = Arc::new(BudgetState::new());
let config = BudgetConfig {
global_limit: Some(1000.0),
enforcement: Enforcement::Hard,
..Default::default()
};
let layer = BudgetLayer::new(config, Arc::clone(&state));
let mut svc = layer.layer(StreamingUsageService {
prompt_tokens: 1000,
completion_tokens: 500,
});
let resp = svc
.call(LlmRequest::ChatStream(chat_req("gpt-4")))
.await
.expect("streamed call should succeed");
let LlmResponse::ChatStream(stream) = resp else {
panic!("expected a ChatStream response");
};
drop(stream);
assert!(
state.global_spend() > 0.0,
"spend must be recorded for a stream that was never polled"
);
}
}