Skip to main content

lean_ctx/core/
agent_budget.rs

1use 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
103/// Remove an agent's budget entry entirely. Safe only for agents that can no longer
104/// issue reads (finished / dead PID) — a live agent would have its budget silently
105/// reset to 0 on the next check. Bounds the BUDGETS map on long-lived daemons.
106pub 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}