use std::path::Path;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use anyhow::Result;
use rusqlite::{params, Connection};
const RUNTIME: &str = "claude-code";
const READ_ATTEMPTS: usize = 5;
const READ_RETRY_DELAY: Duration = Duration::from_millis(100);
#[derive(Debug, Default, PartialEq)]
struct TurnUsage {
input_tok: i64,
output_tok: i64,
cache_creation_tok: i64,
cache_read_tok: i64,
usd: f64,
}
pub fn run(root: &Path, agent: &str) -> Result<()> {
let stdin = std::io::read_to_string(std::io::stdin()).unwrap_or_default();
if stdin.trim().is_empty() {
tracing::debug!(agent, "budget-record: empty stdin, nothing to record");
return Ok(());
}
let transcript_path = match serde_json::from_str::<serde_json::Value>(&stdin)
.ok()
.and_then(|v| {
v.get("transcript_path")
.and_then(|p| p.as_str())
.map(str::to_string)
}) {
Some(p) => p,
None => {
tracing::debug!(agent, "budget-record: no transcript_path in hook payload");
return Ok(());
}
};
let jsonl = match read_transcript(Path::new(&transcript_path)) {
Some(j) => j,
None => {
tracing::warn!(
agent,
transcript_path,
"budget-record: transcript unreadable, skipping"
);
return Ok(());
}
};
let usage = sum_last_turn_usage(&jsonl);
let compose = super::load(root)?;
let db_path = compose.root.join(&compose.global.broker.path);
if let Some(parent) = db_path.parent() {
std::fs::create_dir_all(parent).ok();
}
let conn = Connection::open(&db_path)?;
conn.busy_timeout(Duration::from_secs(5))?;
conn.pragma_update(None, "journal_mode", "WAL")?;
team_core::mailbox::ensure(&conn)?;
let project_id = agent.split(':').next().unwrap_or(agent);
record_cost(&conn, project_id, agent, &usage, now())?;
tracing::debug!(agent, usd = usage.usd, "budget-record: recorded cost");
Ok(())
}
fn read_transcript(path: &Path) -> Option<String> {
for attempt in 0..READ_ATTEMPTS {
if let Ok(content) = std::fs::read_to_string(path) {
if transcript_is_complete(&content) {
return Some(content);
}
}
if attempt + 1 < READ_ATTEMPTS {
std::thread::sleep(READ_RETRY_DELAY);
}
}
std::fs::read_to_string(path).ok()
}
fn transcript_is_complete(content: &str) -> bool {
match content.lines().rfind(|l| !l.trim().is_empty()) {
Some(last) => serde_json::from_str::<serde_json::Value>(last).is_ok(),
None => false,
}
}
fn sum_last_turn_usage(jsonl: &str) -> TurnUsage {
let entries: Vec<serde_json::Value> = jsonl
.lines()
.filter(|l| !l.trim().is_empty())
.filter_map(|l| serde_json::from_str::<serde_json::Value>(l).ok())
.collect();
let entry_type = |v: &serde_json::Value| {
v.get("type")
.and_then(|t| t.as_str())
.unwrap_or_default()
.to_string()
};
let is_real_prompt = |v: &serde_json::Value| {
if entry_type(v) != "user" {
return false;
}
match v.get("message").and_then(|m| m.get("content")) {
Some(serde_json::Value::Array(blocks)) => !blocks
.iter()
.any(|b| b.get("type").and_then(|t| t.as_str()) == Some("tool_result")),
_ => true,
}
};
let turn_start = entries
.iter()
.rposition(is_real_prompt)
.map_or(0, |i| i + 1);
let mut usage = TurnUsage::default();
let mut seen_ids = std::collections::HashSet::new();
for entry in &entries[turn_start..] {
if entry_type(entry) != "assistant" {
continue;
}
let message = match entry.get("message") {
Some(m) => m,
None => continue,
};
if let Some(id) = message.get("id").and_then(|v| v.as_str()) {
if !seen_ids.insert(id.to_string()) {
continue;
}
}
let model = message
.get("model")
.and_then(|m| m.as_str())
.unwrap_or_default();
let u = message.get("usage");
let tok = |key: &str| -> i64 {
u.and_then(|u| u.get(key))
.and_then(serde_json::Value::as_i64)
.unwrap_or(0)
};
let input = tok("input_tokens");
let output = tok("output_tokens");
let cache_creation = tok("cache_creation_input_tokens");
let cache_read = tok("cache_read_input_tokens");
usage.input_tok += input;
usage.output_tok += output;
usage.cache_creation_tok += cache_creation;
usage.cache_read_tok += cache_read;
usage.usd += price_usd(model, input, output, cache_creation, cache_read);
}
usage
}
fn price_usd(
model: &str,
input_tok: i64,
output_tok: i64,
cache_creation_tok: i64,
cache_read_tok: i64,
) -> f64 {
let (input_rate, output_rate) = if model.contains("opus") {
(15.0, 75.0)
} else if model.contains("sonnet") {
(3.0, 15.0)
} else if model.contains("haiku") {
(1.0, 5.0)
} else {
(0.0, 0.0)
};
let per_million = |tokens: i64, rate: f64| (tokens as f64) * rate / 1_000_000.0;
per_million(input_tok, input_rate)
+ per_million(output_tok, output_rate)
+ per_million(cache_read_tok, input_rate * 0.1)
+ per_million(cache_creation_tok, input_rate * 1.25)
}
fn record_cost(
conn: &Connection,
project_id: &str,
agent_id: &str,
usage: &TurnUsage,
observed_at: f64,
) -> Result<()> {
let input_tok = usage.input_tok + usage.cache_creation_tok + usage.cache_read_tok;
conn.execute(
"INSERT INTO budget (project_id, agent_id, runtime, usd, input_tok, output_tok, observed_at)
VALUES (?1,?2,?3,?4,?5,?6,?7)",
params![
project_id,
agent_id,
RUNTIME,
usage.usd,
input_tok,
usage.output_tok,
observed_at
],
)?;
Ok(())
}
fn now() -> f64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs_f64())
.unwrap_or(0.0)
}
#[cfg(test)]
mod tests {
use super::*;
fn db() -> Connection {
let conn = Connection::open_in_memory().expect("open in-memory db");
team_core::mailbox::ensure(&conn).expect("bootstrap schema");
conn
}
#[test]
fn sum_last_turn_usage_sums_only_the_current_turn() {
let jsonl = r#"
{"type":"user","message":{"role":"user","content":"old"}}
{"type":"assistant","message":{"model":"claude-opus-4-8","usage":{"input_tokens":999,"output_tokens":999}}}
{"type":"user","message":{"role":"user","content":"current"}}
{"type":"assistant","message":{"model":"claude-sonnet-4-6","usage":{"input_tokens":100,"output_tokens":50,"cache_creation_input_tokens":10,"cache_read_input_tokens":200}}}
{"type":"user","message":{"role":"user","content":[{"type":"tool_result","tool_use_id":"x","content":"ok"}]}}
{"type":"assistant","message":{"model":"claude-sonnet-4-6","usage":{"input_tokens":40,"output_tokens":20}}}
"#;
let usage = sum_last_turn_usage(jsonl);
assert_eq!(
usage.input_tok, 140,
"100 + 40: both current-turn assistant msgs, tool-result is not a boundary"
);
assert_eq!(usage.output_tok, 70, "50 + 20 from the current turn only");
assert_eq!(usage.cache_creation_tok, 10);
assert_eq!(usage.cache_read_tok, 200);
let expected = price_usd("claude-sonnet-4-6", 100, 50, 10, 200)
+ price_usd("claude-sonnet-4-6", 40, 20, 0, 0);
assert!(
(usage.usd - expected).abs() < 1e-9,
"usd should sum each assistant message priced by its own model: {} vs {expected}",
usage.usd
);
}
#[test]
fn sum_last_turn_usage_tolerates_missing_fields_and_bad_lines() {
let jsonl = r#"
{"type":"user","message":{"role":"user"}}
{"type":"assistant","message":{"model":"claude-opus-4-8"}}
{"type":"assistant","message":{"model":"claude-opus-4-8","usage":{"input_tokens":10}}}
{not valid json
"#;
let usage = sum_last_turn_usage(jsonl);
assert_eq!(usage.input_tok, 10);
assert_eq!(usage.output_tok, 0);
assert_eq!(usage.cache_creation_tok, 0);
assert_eq!(usage.cache_read_tok, 0);
}
#[test]
fn sum_last_turn_usage_dedups_repeated_message_id() {
let jsonl = r#"
{"type":"user","message":{"role":"user","content":"go"}}
{"type":"assistant","message":{"id":"msg_A","model":"claude-opus-4-8","usage":{"input_tokens":100,"output_tokens":50}}}
{"type":"assistant","message":{"id":"msg_A","model":"claude-opus-4-8","usage":{"input_tokens":100,"output_tokens":50}}}
{"type":"assistant","message":{"id":"msg_B","model":"claude-opus-4-8","usage":{"input_tokens":30,"output_tokens":10}}}
"#;
let usage = sum_last_turn_usage(jsonl);
assert_eq!(
usage.input_tok, 130,
"msg_A (100) counted once + msg_B (30), not msg_A twice"
);
assert_eq!(usage.output_tok, 60, "msg_A (50) once + msg_B (10)");
}
#[test]
fn price_usd_matches_each_tier_by_substring() {
assert!((price_usd("claude-opus-4-8", 1_000_000, 1_000_000, 0, 0) - 90.0).abs() < 1e-9);
assert!((price_usd("claude-sonnet-4-6", 1_000_000, 1_000_000, 0, 0) - 18.0).abs() < 1e-9);
assert!((price_usd("claude-3-5-haiku", 1_000_000, 1_000_000, 0, 0) - 6.0).abs() < 1e-9);
}
#[test]
fn price_usd_unknown_model_is_zero() {
assert_eq!(
price_usd("gpt-5", 1_000_000, 1_000_000, 1_000_000, 1_000_000),
0.0
);
}
#[test]
fn price_usd_prices_cache_tokens_off_the_input_rate() {
let usd = price_usd("claude-opus-4-8", 0, 0, 1_000_000, 1_000_000);
assert!(
(usd - 20.25).abs() < 1e-9,
"cache pricing: 1M creation (18.75) + 1M read (1.5): {usd}"
);
}
#[test]
fn record_cost_inserts_one_row_with_summed_tokens() {
let conn = db();
let usage = TurnUsage {
input_tok: 100,
output_tok: 50,
cache_creation_tok: 10,
cache_read_tok: 200,
usd: 1.23,
};
record_cost(&conn, "alpha", "alpha:dev", &usage, 456.0).expect("record cost");
let (project_id, agent_id, runtime, usd, input_tok, output_tok, observed_at): (
String,
String,
String,
f64,
i64,
i64,
f64,
) = conn
.query_row(
"SELECT project_id, agent_id, runtime, usd, input_tok, output_tok, observed_at \
FROM budget",
[],
|r| {
Ok((
r.get(0)?,
r.get(1)?,
r.get(2)?,
r.get(3)?,
r.get(4)?,
r.get(5)?,
r.get(6)?,
))
},
)
.expect("query the inserted row");
assert_eq!(project_id, "alpha");
assert_eq!(agent_id, "alpha:dev");
assert_eq!(runtime, "claude-code");
assert!((usd - 1.23).abs() < 1e-9);
assert_eq!(input_tok, 310);
assert_eq!(output_tok, 50);
assert_eq!(observed_at, 456.0);
}
}