lean_ctx/core/
agent_budget.rs1use std::collections::HashMap;
2use std::sync::Mutex;
3
4use serde::{Deserialize, Serialize};
5
6static BUDGETS: Mutex<Option<HashMap<String, AgentBudget>>> = Mutex::new(None);
7
8#[derive(Debug, Clone, Serialize, Deserialize)]
9pub struct AgentBudget {
10 pub agent_id: String,
11 pub token_limit: usize,
12 pub tokens_consumed: usize,
13 pub reads_count: u32,
14 pub last_reset: String,
15}
16
17#[derive(Debug, Clone, PartialEq)]
18pub enum BudgetCheckResult {
19 Allowed { remaining: usize },
20 Exceeded { limit: usize, consumed: usize },
21 Warning { remaining: usize, percent_used: f32 },
22}
23
24const WARNING_THRESHOLD: f32 = 0.80;
25
26fn with_budgets<F, R>(f: F) -> R
27where
28 F: FnOnce(&mut HashMap<String, AgentBudget>) -> R,
29{
30 let mut guard = BUDGETS
31 .lock()
32 .unwrap_or_else(std::sync::PoisonError::into_inner);
33 let map = guard.get_or_insert_with(HashMap::new);
34 f(map)
35}
36
37fn ensure_entry<'a>(
38 map: &'a mut HashMap<String, AgentBudget>,
39 agent_id: &str,
40) -> &'a mut AgentBudget {
41 map.entry(agent_id.to_string())
42 .or_insert_with(|| AgentBudget {
43 agent_id: agent_id.to_string(),
44 token_limit: usize::MAX,
45 tokens_consumed: 0,
46 reads_count: 0,
47 last_reset: chrono::Utc::now().to_rfc3339(),
48 })
49}
50
51pub fn check_budget(agent_id: &str, tokens_to_consume: usize) -> BudgetCheckResult {
52 with_budgets(|map| {
53 let budget = ensure_entry(map, agent_id);
54 if budget.token_limit == usize::MAX || budget.token_limit == 0 {
55 return BudgetCheckResult::Allowed {
56 remaining: usize::MAX,
57 };
58 }
59
60 let projected = budget.tokens_consumed.saturating_add(tokens_to_consume);
61 if projected > budget.token_limit {
62 return BudgetCheckResult::Exceeded {
63 limit: budget.token_limit,
64 consumed: budget.tokens_consumed,
65 };
66 }
67
68 let percent_used = projected as f32 / budget.token_limit as f32;
69 let remaining = budget.token_limit.saturating_sub(projected);
70
71 if percent_used >= WARNING_THRESHOLD {
72 BudgetCheckResult::Warning {
73 remaining,
74 percent_used,
75 }
76 } else {
77 BudgetCheckResult::Allowed { remaining }
78 }
79 })
80}
81
82pub fn record_consumption(agent_id: &str, tokens: usize) {
83 with_budgets(|map| {
84 let budget = ensure_entry(map, agent_id);
85 budget.tokens_consumed = budget.tokens_consumed.saturating_add(tokens);
86 budget.reads_count += 1;
87 });
88}
89
90pub fn get_status(agent_id: &str) -> AgentBudget {
91 with_budgets(|map| ensure_entry(map, agent_id).clone())
92}
93
94pub fn reset(agent_id: &str) {
95 with_budgets(|map| {
96 let budget = ensure_entry(map, agent_id);
97 budget.tokens_consumed = 0;
98 budget.reads_count = 0;
99 budget.last_reset = chrono::Utc::now().to_rfc3339();
100 });
101}
102
103pub fn remove(agent_id: &str) {
107 with_budgets(|map| {
108 map.remove(agent_id);
109 });
110}
111
112pub fn set_limit(agent_id: &str, limit: usize) {
113 with_budgets(|map| {
114 let budget = ensure_entry(map, agent_id);
115 budget.token_limit = if limit == 0 { usize::MAX } else { limit };
116 });
117}
118
119pub fn init_from_config() {
120 let cfg_limit = crate::core::config::Config::load().agent_token_budget;
121 if cfg_limit > 0 {
122 with_budgets(|map| {
123 for budget in map.values_mut() {
124 if budget.token_limit == usize::MAX {
125 budget.token_limit = cfg_limit;
126 }
127 }
128 });
129 }
130}
131
132pub fn default_limit_from_config() -> usize {
133 let cfg_limit = crate::core::config::Config::load().agent_token_budget;
134 if cfg_limit == 0 {
135 usize::MAX
136 } else {
137 cfg_limit
138 }
139}
140
141#[cfg(test)]
142mod tests {
143 use super::*;
144
145 fn test_agent(name: &str) -> String {
146 format!("test_agent_{name}_{:?}", std::thread::current().id())
147 }
148
149 #[test]
150 fn unlimited_budget_always_allows() {
151 let id = test_agent("unlimited");
152 let result = check_budget(&id, 1_000_000);
153 assert!(matches!(result, BudgetCheckResult::Allowed { .. }));
154 }
155
156 #[test]
157 fn set_limit_and_exceed() {
158 let id = test_agent("exceed");
159 set_limit(&id, 1000);
160 record_consumption(&id, 800);
161 let result = check_budget(&id, 300);
162 assert!(matches!(
163 result,
164 BudgetCheckResult::Exceeded {
165 limit: 1000,
166 consumed: 800
167 }
168 ));
169 }
170
171 #[test]
172 fn warning_at_80_percent() {
173 let id = test_agent("warning");
174 set_limit(&id, 1000);
175 record_consumption(&id, 700);
176 let result = check_budget(&id, 100);
177 assert!(matches!(result, BudgetCheckResult::Warning { .. }));
178 }
179
180 #[test]
181 fn reset_clears_consumption() {
182 let id = test_agent("reset");
183 set_limit(&id, 1000);
184 record_consumption(&id, 900);
185 reset(&id);
186 let status = get_status(&id);
187 assert_eq!(status.tokens_consumed, 0);
188 assert_eq!(status.reads_count, 0);
189 }
190
191 #[test]
192 fn zero_limit_means_unlimited() {
193 let id = test_agent("zero");
194 set_limit(&id, 0);
195 let result = check_budget(&id, 1_000_000);
196 assert!(matches!(result, BudgetCheckResult::Allowed { .. }));
197 }
198
199 #[test]
200 fn record_increments_reads_count() {
201 let id = test_agent("reads");
202 record_consumption(&id, 100);
203 record_consumption(&id, 200);
204 let status = get_status(&id);
205 assert_eq!(status.reads_count, 2);
206 assert_eq!(status.tokens_consumed, 300);
207 }
208}