Skip to main content

atman_runtime/
cost.rs

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}