Skip to main content

aether_core/context/
session_usage_tracker.rs

1use llm::{LlmCallPurpose, ModelIdentity, SessionUsageEvent, SessionUsageTotals, TokenUsage, UsageCost, UsageSource};
2
3/// Billable usage for one agent session. The root agent's session tracker captures totals for all agents (e.g. sub-agents).
4#[derive(Clone, Debug)]
5pub struct SessionUsageTracker {
6    source: UsageSource,
7    sequence: u64,
8    totals: SessionUsageTotals,
9}
10
11impl SessionUsageTracker {
12    pub fn new(agent_name: impl Into<String>) -> Self {
13        Self { source: UsageSource::new(agent_name), sequence: 0, totals: SessionUsageTotals::default() }
14    }
15
16    /// Continue from the last persisted event so a resumed session keeps its totals.
17    pub fn resume_from(&mut self, last: &SessionUsageEvent) {
18        self.sequence = last.sequence;
19        self.totals = last.totals.clone();
20    }
21
22    pub fn source(&self) -> &UsageSource {
23        &self.source
24    }
25
26    /// Record one of this agent's own calls, priced from the model's catalog entry.
27    pub fn record(&mut self, purpose: LlmCallPurpose, model: ModelIdentity, tokens: TokenUsage) -> SessionUsageEvent {
28        let estimated_cost = model.pricing.map(|pricing| pricing.estimate_cost(tokens));
29        self.push(self.source.clone(), purpose, model, tokens, estimated_cost)
30    }
31
32    /// Fold a sub-agent's sample into these totals. The child's model and cost
33    /// stand; its lineage is filled in when the child did not know it.
34    pub fn record_child(&mut self, task_id: &str, child: SessionUsageEvent) -> SessionUsageEvent {
35        let UsageSource { agent_id, parent_agent_id, task_id: child_task_id, agent_name } = child.source;
36        let source = UsageSource {
37            agent_id,
38            parent_agent_id: parent_agent_id.or_else(|| Some(self.source.agent_id.clone())),
39            task_id: child_task_id.or_else(|| Some(task_id.to_string())),
40            agent_name,
41        };
42        self.push(source, child.purpose, child.model, child.tokens, child.estimated_cost)
43    }
44
45    fn push(
46        &mut self,
47        source: UsageSource,
48        purpose: LlmCallPurpose,
49        model: ModelIdentity,
50        tokens: TokenUsage,
51        estimated_cost: Option<UsageCost>,
52    ) -> SessionUsageEvent {
53        self.sequence = self.sequence.saturating_add(1);
54        self.totals.add(tokens, estimated_cost);
55        SessionUsageEvent {
56            sequence: self.sequence,
57            source,
58            purpose,
59            model,
60            tokens,
61            estimated_cost,
62            totals: self.totals.clone(),
63        }
64    }
65}
66
67#[cfg(test)]
68mod tests {
69    use super::*;
70    use llm::testing::{priced_model, session_usage_event};
71
72    #[test]
73    fn own_calls_are_sequenced_priced_and_totalled() {
74        let mut tracker = SessionUsageTracker::new("root");
75        let first = tracker.record(LlmCallPurpose::Chat, ModelIdentity::default(), TokenUsage::new(2, 3));
76        let second =
77            tracker.record(LlmCallPurpose::Chat, ModelIdentity::of(Some(&priced_model())), TokenUsage::new(5, 7));
78
79        assert_eq!((first.sequence, second.sequence), (1, 2));
80        assert_eq!(first.source, *tracker.source());
81        assert!(first.estimated_cost.is_none());
82        assert!(second.estimated_cost.is_some());
83        assert_eq!(second.totals.tokens.input_tokens.get(), 7);
84        assert_eq!(second.totals.tokens.output_tokens.get(), 10);
85        assert_eq!(second.totals.unpriced_calls, 1);
86        assert!(second.totals.estimated_usd.get() > 0.0);
87    }
88
89    #[test]
90    fn zero_token_samples_without_pricing_stay_fully_priced() {
91        let mut tracker = SessionUsageTracker::new("root");
92        let event = tracker.record(LlmCallPurpose::Chat, ModelIdentity::default(), TokenUsage::default());
93        assert!(event.totals.is_fully_priced());
94    }
95
96    #[test]
97    fn resumed_tracker_continues_sequence_and_totals() {
98        let mut tracker = SessionUsageTracker::new("root");
99        let last = tracker.record(LlmCallPurpose::Chat, ModelIdentity::default(), TokenUsage::new(4, 6));
100
101        let mut resumed = SessionUsageTracker::new("root");
102        resumed.resume_from(&last);
103        let next = resumed.record(LlmCallPurpose::Compaction, ModelIdentity::default(), TokenUsage::new(1, 1));
104        assert_eq!(next.sequence, 2);
105        assert_eq!(next.totals.tokens.input_tokens.get(), 5);
106        assert_eq!(next.totals.tokens.output_tokens.get(), 7);
107        assert_eq!(next.totals.unpriced_calls, 2);
108    }
109
110    #[test]
111    fn child_samples_are_resequenced_totalled_and_given_lineage() {
112        let mut tracker = SessionUsageTracker::new("root");
113        let mut child = session_usage_event(9, TokenUsage::new(8, 4));
114        child.source = UsageSource::new("explorer");
115
116        let folded = tracker.record_child("task_0", child.clone());
117        let own = tracker.record(LlmCallPurpose::Chat, ModelIdentity::default(), TokenUsage::new(1, 1));
118
119        assert_eq!(folded.sequence, 1);
120        assert_eq!(folded.source.agent_id, child.source.agent_id);
121        assert_eq!(folded.source.agent_name, "explorer");
122        assert_eq!(folded.source.parent_agent_id.as_deref(), Some(tracker.source().agent_id.as_str()));
123        assert_eq!(folded.source.task_id.as_deref(), Some("task_0"));
124        assert_eq!(folded.tokens, child.tokens);
125        assert_eq!(own.sequence, 2);
126        assert_eq!(own.totals.tokens.input_tokens.get(), 9);
127        assert_eq!(own.totals.unpriced_calls, 2);
128    }
129
130    #[test]
131    fn child_samples_keep_lineage_they_already_carry() {
132        let mut tracker = SessionUsageTracker::new("root");
133        let mut grandchild = session_usage_event(1, TokenUsage::new(1, 1));
134        grandchild.source.parent_agent_id = Some("middle".to_string());
135        grandchild.source.task_id = Some("task_7".to_string());
136
137        let folded = tracker.record_child("task_0", grandchild);
138        assert_eq!(folded.source.parent_agent_id.as_deref(), Some("middle"));
139        assert_eq!(folded.source.task_id.as_deref(), Some("task_7"));
140    }
141}