use std::collections::HashMap;
use std::sync::Arc;
use parking_lot::Mutex;
use serde::{Deserialize, Serialize};
use thiserror::Error;
use tracing::warn;
#[derive(Debug, Error)]
pub enum BudgetError {
#[error(
"budget exceeded for agent={agent} tool={tool}: spent ${spent_usd:.4} of ${cap_usd:.4}"
)]
Exceeded {
agent: String,
tool: String,
spent_usd: f64,
cap_usd: f64,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolBudget {
pub cap_usd: f64,
pub window_secs: u64,
}
impl Default for ToolBudget {
fn default() -> Self {
Self {
cap_usd: 1.0,
window_secs: 3600,
}
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ToolSpend {
pub spent_usd: f64,
pub calls: u64,
pub denied: u64,
pub window_started_at_unix: i64,
}
pub struct McpBudgetGateway {
budgets: Mutex<HashMap<(String, String), ToolBudget>>,
spend: Mutex<HashMap<(String, String), ToolSpend>>,
default_budget: ToolBudget,
}
impl McpBudgetGateway {
pub fn new(default_budget: ToolBudget) -> Arc<Self> {
Arc::new(Self {
budgets: Mutex::new(HashMap::new()),
spend: Mutex::new(HashMap::new()),
default_budget,
})
}
pub fn set_budget(&self, agent: &str, tool: &str, b: ToolBudget) {
self.budgets
.lock()
.insert((agent.to_string(), tool.to_string()), b);
}
pub fn check(&self, agent: &str, tool: &str) -> Result<ToolSpend, BudgetError> {
let key = (agent.to_string(), tool.to_string());
let budget = self
.budgets
.lock()
.get(&key)
.cloned()
.unwrap_or_else(|| self.default_budget.clone());
let mut spend_map = self.spend.lock();
let now = now_unix();
let spend = spend_map.entry(key.clone()).or_default();
if budget.window_secs > 0
&& spend.window_started_at_unix > 0
&& (now - spend.window_started_at_unix) as u64 >= budget.window_secs
{
*spend = ToolSpend {
window_started_at_unix: now,
..Default::default()
};
}
if spend.window_started_at_unix == 0 {
spend.window_started_at_unix = now;
}
if spend.spent_usd >= budget.cap_usd {
spend.denied += 1;
let denied = spend.clone();
warn!(
agent = %agent,
tool = %tool,
spent = spend.spent_usd,
cap = budget.cap_usd,
"MCP budget denied"
);
return Err(BudgetError::Exceeded {
agent: agent.into(),
tool: tool.into(),
spent_usd: denied.spent_usd,
cap_usd: budget.cap_usd,
});
}
Ok(spend.clone())
}
pub fn record(&self, agent: &str, tool: &str, cost_usd: f64) {
let key = (agent.to_string(), tool.to_string());
let mut spend_map = self.spend.lock();
let spend = spend_map.entry(key).or_default();
if spend.window_started_at_unix == 0 {
spend.window_started_at_unix = now_unix();
}
spend.spent_usd += cost_usd;
spend.calls += 1;
}
pub fn snapshot(&self) -> HashMap<String, ToolSpend> {
self.spend
.lock()
.iter()
.map(|((agent, tool), s)| (format!("{agent}::{tool}"), s.clone()))
.collect()
}
}
fn now_unix() -> i64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs() as i64)
.unwrap_or(0)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn check_allows_until_cap() {
let g = McpBudgetGateway::new(ToolBudget {
cap_usd: 0.10,
window_secs: 3600,
});
for _ in 0..5 {
g.check("agent-a", "search").unwrap();
g.record("agent-a", "search", 0.02);
}
assert!(g.check("agent-a", "search").is_err());
}
#[test]
fn explicit_budget_overrides_default() {
let g = McpBudgetGateway::new(ToolBudget {
cap_usd: 100.0,
window_secs: 3600,
});
g.set_budget(
"agent-a",
"expensive_tool",
ToolBudget {
cap_usd: 0.05,
window_secs: 3600,
},
);
g.check("agent-a", "expensive_tool").unwrap();
g.record("agent-a", "expensive_tool", 0.06);
assert!(g.check("agent-a", "expensive_tool").is_err());
g.check("agent-a", "cheap_tool").unwrap();
}
#[test]
fn snapshot_includes_calls_and_denials() {
let g = McpBudgetGateway::new(ToolBudget {
cap_usd: 0.01,
window_secs: 3600,
});
g.check("a", "t").unwrap();
g.record("a", "t", 0.02);
let _ = g.check("a", "t"); let snap = g.snapshot();
let entry = snap.get("a::t").unwrap();
assert_eq!(entry.calls, 1);
assert_eq!(entry.denied, 1);
assert!((entry.spent_usd - 0.02).abs() < 1e-9);
}
}