use std::sync::Arc;
use parking_lot::Mutex;
use theway_llm_provider::{Message as PiMessage, Usage};
use crate::types::{AgentMessage, LoopEvent};
#[derive(Clone, Debug, Default)]
pub struct CostSnapshot {
pub tokens: Usage,
pub turn_count: u64,
}
impl CostSnapshot {
pub fn total_cost(&self) -> f64 {
self.tokens.cost.total
}
}
#[derive(Clone, Debug)]
pub struct CostTracker {
inner: Arc<Mutex<CostSnapshot>>,
}
impl Default for CostTracker {
fn default() -> Self {
Self::new()
}
}
impl CostTracker {
pub fn new() -> Self {
Self {
inner: Arc::new(Mutex::new(CostSnapshot::default())),
}
}
pub fn snapshot(&self) -> CostSnapshot {
self.inner.lock().clone()
}
pub fn reset(&self) {
*self.inner.lock() = CostSnapshot::default();
}
pub fn record(&self, usage: &Usage) {
let mut g = self.inner.lock();
g.tokens.input = g.tokens.input.saturating_add(usage.input);
g.tokens.output = g.tokens.output.saturating_add(usage.output);
g.tokens.cache_read = g.tokens.cache_read.saturating_add(usage.cache_read);
g.tokens.cache_write = g.tokens.cache_write.saturating_add(usage.cache_write);
g.tokens.total_tokens = g.tokens.total_tokens.saturating_add(usage.total_tokens);
let c = &mut g.tokens.cost;
c.input += usage.cost.input;
c.output += usage.cost.output;
c.cache_read += usage.cost.cache_read;
c.cache_write += usage.cost.cache_write;
c.total += usage.cost.total;
g.turn_count += 1;
}
pub fn as_callback(&self) -> crate::agent::LoopSyncCallback {
let tracker = self.clone();
Arc::new(move |event| {
if let LoopEvent::MessageEnd {
message: AgentMessage::Llm(PiMessage::Assistant(a)),
} = event
{
tracker.record(&a.usage);
}
})
}
}
pub fn one_line_summary(snap: &CostSnapshot) -> String {
format!(
"tokens: in={} out={} cached={} total={} | cost ${:.4}",
snap.tokens.input,
snap.tokens.output,
snap.tokens.cache_read + snap.tokens.cache_write,
snap.tokens.total_tokens,
snap.total_cost()
)
}
pub fn full_breakdown(snap: &CostSnapshot) -> String {
let c = &snap.tokens.cost;
format!(
" turns: {turns}\n\
\n\
Tokens:\n\
\n input {input}\n output {output}\n cache read {cache_read}\n cache write {cache_write}\n total {total}\n\n\
Cost (USD):\n\
\n input ${ci:.4}\n output ${co:.4}\n cache read ${cr:.4}\n cache write ${cw:.4}\n total ${ct:.4}\n",
turns = snap.turn_count,
input = snap.tokens.input,
output = snap.tokens.output,
cache_read = snap.tokens.cache_read,
cache_write = snap.tokens.cache_write,
total = snap.tokens.total_tokens,
ci = c.input,
co = c.output,
cr = c.cache_read,
cw = c.cache_write,
ct = c.total,
)
}
#[cfg(test)]
mod cost_bridge_tests {
tests_bridge_macro::tests_bridge!("agent/cost");
}