1use llm::{LlmCallPurpose, ModelIdentity, SessionUsageEvent, SessionUsageTotals, TokenUsage, UsageCost, UsageSource};
2
3#[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 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 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 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 nested_samples_and_resumption_preserve_cumulative_cost_components() {
132 let model = ModelIdentity::of(Some(&priced_model()));
133 let tokens = TokenUsage {
134 cache_read_tokens: Some(20.into()),
135 cache_creation_tokens: Some(10.into()),
136 ..TokenUsage::new(100, 50)
137 };
138
139 let mut root = SessionUsageTracker::new("root");
140 root.record(LlmCallPurpose::Chat, model.clone(), tokens);
141
142 let mut child = SessionUsageTracker::new("child");
143 let child_sample = child.record(LlmCallPurpose::Chat, model.clone(), tokens);
144 root.record_child("child-task", child_sample);
145
146 let mut grandchild = SessionUsageTracker::new("grandchild");
147 let sample = grandchild.record(LlmCallPurpose::Compaction, model.clone(), tokens);
148 let sample_cost = sample.estimated_cost.unwrap();
149 let folded = child.record_child("grandchild-task", sample);
150
151 let last = root.record_child("child-task", folded);
152
153 assert_eq!(last.sequence, 3);
154 assert_eq!(last.source.agent_id, grandchild.source().agent_id);
155 assert_eq!(last.totals.tokens.input_tokens.get(), 300);
156 assert_eq!(last.estimated_cost, Some(sample_cost));
157 assert_cost_components(&last, sample_cost, 3);
158
159 let persisted = serde_json::to_string(&last).unwrap();
160 let mut resumed = SessionUsageTracker::new("root");
161 resumed.resume_from(&serde_json::from_str(&persisted).unwrap());
162 let next = resumed.record(LlmCallPurpose::Chat, model, tokens);
163 assert_eq!(next.sequence, 4);
164 assert_eq!(next.totals.tokens.input_tokens.get(), 400);
165 assert_cost_components(&next, sample_cost, 4);
166 }
167
168 fn assert_cost_components(event: &SessionUsageEvent, cost: UsageCost, samples: usize) {
169 let totals = serde_json::to_value(&event.totals).unwrap();
170 for (field, amount) in [
171 ("estimated_input_usd", cost.input_usd),
172 ("estimated_output_usd", cost.output_usd),
173 ("estimated_cache_read_usd", cost.cache_read_usd),
174 ("estimated_cache_creation_usd", cost.cache_creation_usd),
175 ("estimated_usd", cost.total_usd),
176 ] {
177 let expected = (0..samples).fold(llm::Usd::ZERO, |total, _| total + amount);
178 assert_eq!(totals[field], serde_json::json!(expected), "{field}");
179 }
180 }
181
182 #[test]
183 fn child_samples_keep_lineage_they_already_carry() {
184 let mut tracker = SessionUsageTracker::new("root");
185 let mut grandchild = session_usage_event(1, TokenUsage::new(1, 1));
186 grandchild.source.parent_agent_id = Some("middle".to_string());
187 grandchild.source.task_id = Some("task_7".to_string());
188
189 let folded = tracker.record_child("task_0", grandchild);
190 assert_eq!(folded.source.parent_agent_id.as_deref(), Some("middle"));
191 assert_eq!(folded.source.task_id.as_deref(), Some("task_7"));
192 }
193}