1use std::collections::HashMap;
2
3use crate::event::{Event, LlmCallStatus};
4use crate::provider::TokenUsage;
5
6#[derive(Debug, Default, Clone)]
7pub struct CostSummary {
8 pub calls: u64,
9 pub failures: u64,
10 pub usage: TokenUsage,
11 pub wallclock_ms: u64,
12}
13
14impl CostSummary {
15 fn accumulate(&mut self, usage: &TokenUsage, wallclock_ms: u64, failed: bool) {
16 self.calls += 1;
17 if failed {
18 self.failures += 1;
19 }
20 self.usage.input = self.usage.input.saturating_add(usage.input);
21 self.usage.cached_input = self.usage.cached_input.saturating_add(usage.cached_input);
22 self.usage.output = self.usage.output.saturating_add(usage.output);
23 self.usage.cache_write = self.usage.cache_write.saturating_add(usage.cache_write);
24 self.usage.reasoning_tokens = self
25 .usage
26 .reasoning_tokens
27 .saturating_add(usage.reasoning_tokens);
28 self.wallclock_ms = self.wallclock_ms.saturating_add(wallclock_ms);
29 }
30}
31
32pub fn summarize_by_model(events: &[Event]) -> HashMap<String, CostSummary> {
33 let mut out: HashMap<String, CostSummary> = HashMap::new();
34 for e in events {
35 if let Event::LlmCall {
36 model,
37 usage,
38 wallclock_ms,
39 status,
40 ..
41 } = e
42 {
43 let entry = out.entry(model.clone()).or_default();
44 entry.accumulate(
45 usage,
46 *wallclock_ms,
47 matches!(status, LlmCallStatus::Errored { .. }),
48 );
49 }
50 }
51 out
52}
53
54pub fn summarize_by_provider(events: &[Event]) -> HashMap<String, CostSummary> {
55 let mut out: HashMap<String, CostSummary> = HashMap::new();
56 for e in events {
57 if let Event::LlmCall {
58 provider,
59 usage,
60 wallclock_ms,
61 status,
62 ..
63 } = e
64 {
65 let entry = out.entry(provider.clone()).or_default();
66 entry.accumulate(
67 usage,
68 *wallclock_ms,
69 matches!(status, LlmCallStatus::Errored { .. }),
70 );
71 }
72 }
73 out
74}
75
76pub fn total(events: &[Event]) -> CostSummary {
77 let mut acc = CostSummary::default();
78 for e in events {
79 if let Event::LlmCall {
80 usage,
81 wallclock_ms,
82 status,
83 ..
84 } = e
85 {
86 acc.accumulate(
87 usage,
88 *wallclock_ms,
89 matches!(status, LlmCallStatus::Errored { .. }),
90 );
91 }
92 }
93 acc
94}
95
96#[cfg(test)]
97mod tests {
98 use super::*;
99
100 fn ok_call(model: &str, provider: &str, in_tok: u64, out_tok: u64) -> Event {
101 Event::LlmCall {
102 model: model.into(),
103 provider: provider.into(),
104 context_plan_id: None,
105 context_epoch: None,
106 context_tokens: None,
107 usage_source: None,
108 context_call_purpose: None,
109 context_call_identity: None,
110 context_cache: None,
111 assistant_tool_batch_width: None,
112 usage: TokenUsage {
113 input: in_tok,
114 output: out_tok,
115 ..Default::default()
116 },
117 wallclock_ms: 100,
118 status: LlmCallStatus::Ok,
119 ttft_ms: None,
120 tokens_per_second: None,
121 run_id: None,
122 node_id: None,
123 }
124 }
125
126 fn err_call(model: &str) -> Event {
127 Event::LlmCall {
128 model: model.into(),
129 provider: "p".into(),
130 context_plan_id: None,
131 context_epoch: None,
132 context_tokens: None,
133 usage_source: None,
134 context_call_purpose: None,
135 context_call_identity: None,
136 context_cache: None,
137 assistant_tool_batch_width: None,
138 usage: TokenUsage {
139 input: 5,
140 ..Default::default()
141 },
142 wallclock_ms: 50,
143 status: LlmCallStatus::Errored {
144 message: "boom".into(),
145 },
146 ttft_ms: None,
147 tokens_per_second: None,
148 run_id: None,
149 node_id: None,
150 }
151 }
152
153 #[test]
154 fn summarize_by_model_groups_multiple_calls() {
155 let events = vec![
156 ok_call("opus", "anthropic", 10, 20),
157 ok_call("opus", "anthropic", 5, 8),
158 ok_call("mini", "openai", 3, 4),
159 ];
160 let s = summarize_by_model(&events);
161 assert_eq!(s.get("opus").unwrap().calls, 2);
162 assert_eq!(s.get("opus").unwrap().usage.input, 15);
163 assert_eq!(s.get("opus").unwrap().usage.output, 28);
164 assert_eq!(s.get("mini").unwrap().calls, 1);
165 }
166
167 #[test]
168 fn failures_counted_separately() {
169 let events = vec![ok_call("m", "p", 1, 2), err_call("m"), err_call("m")];
170 let s = summarize_by_model(&events);
171 assert_eq!(s.get("m").unwrap().calls, 3);
172 assert_eq!(s.get("m").unwrap().failures, 2);
173 }
174
175 #[test]
176 fn total_sums_wallclock() {
177 let events = vec![ok_call("a", "p", 1, 1), ok_call("b", "p", 2, 2)];
178 let t = total(&events);
179 assert_eq!(t.calls, 2);
180 assert_eq!(t.wallclock_ms, 200);
181 }
182
183 #[test]
184 fn total_preserves_every_usage_lane() {
185 let mut event = ok_call("m", "p", 11, 13);
186 let Event::LlmCall { usage, .. } = &mut event else {
187 unreachable!();
188 };
189 usage.cached_input = 17;
190 usage.cache_write = 19;
191 usage.reasoning_tokens = 23;
192
193 let total = total(&[event]);
194
195 assert_eq!(total.usage.input, 11);
196 assert_eq!(total.usage.cached_input, 17);
197 assert_eq!(total.usage.cache_write, 19);
198 assert_eq!(total.usage.output, 13);
199 assert_eq!(total.usage.reasoning_tokens, 23);
200 }
201
202 #[test]
203 fn non_llm_events_ignored() {
204 use crate::event::{FlowRunId, FlowStatus};
205 let run_id = FlowRunId::now();
206 let events = vec![
207 Event::FlowStart {
208 run_id: run_id.clone(),
209 flow_name: "t".into(),
210 parent_run_id: None,
211 parent_node_id: None,
212 spawned: false,
213 },
214 ok_call("m", "p", 5, 5),
215 Event::FlowEnd {
216 run_id,
217 flow_name: "t".into(),
218 status: FlowStatus::Ok,
219 },
220 ];
221 let t = total(&events);
222 assert_eq!(t.calls, 1);
223 assert_eq!(t.usage.input, 5);
224 }
225}