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 dashmap::DashMap;
use tower::{Layer, Service};
use super::types::{LlmRequest, LlmResponse};
use crate::client::BoxFuture;
use crate::cost;
use crate::error::{LiterLlmError, Result};
#[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: 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>>,
}
impl InMemoryBudgetLedger {
#[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,
window,
}
}
#[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 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(&self.per_user, user, ctx.cost_usd, self.window, now);
}
if let Some(key) = ctx.api_key_id {
Self::entry_add(&self.per_api_key, key, ctx.cost_usd, self.window, now);
}
#[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;
if let Some(limit) = self.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) = self.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) = self.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) = self.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) = self.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 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: self.limits.global,
limits_per_user: self.limits.per_user.clone(),
limits_per_api_key: self.limits.per_api_key.clone(),
limits_per_tenant: self.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
}
#[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?;
if let Some(usage) = resp.usage()
&& let Some(usd) = cost::completion_cost(&model, usage.prompt_tokens, usage.completion_tokens)
{
state.record(&model, usd);
if config.enforcement == Enforcement::Soft {
emit_soft_warnings(&config, &state, &model);
}
}
Ok(resp)
})
}
}
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 tower::{Layer as _, Service as _};
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_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"
);
}
}