use super::Error;
use serde_json::Value;
use std::path::Path;
pub(super) const MESSAGES_DIR: &str = "messages";
const TOOL_ORIGIN: &str = "tool";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LastUsage {
pub prompt_tokens: u64,
pub context_window: Option<u64>,
pub model: String,
}
pub(super) fn due(n: Option<u32>, last: Option<&LastUsage>) -> Result<bool, Error> {
let (Some(n), Some(last)) = (n.filter(|n| *n > 0), last) else {
return Ok(false);
};
let Some(window) = last.context_window else {
return Err(Error::CompactionWindowUnknown {
model: last.model.clone(),
});
};
Ok(u128::from(last.prompt_tokens) * 100 >= u128::from(n) * u128::from(window))
}
pub(super) fn last(worktree: &Path) -> Result<Option<LastUsage>, Error> {
let dir = worktree.join(MESSAGES_DIR);
let mut newest: Option<(u32, String, std::path::PathBuf)> = None;
let listing = match std::fs::read_dir(&dir) {
Ok(rd) => rd,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(None),
Err(e) => return Err(Error::Io(e)),
};
for entry in listing {
let path = entry.map_err(Error::Io)?.path();
let Some((seq, model)) = model_entry(&path) else {
continue;
};
if newest.as_ref().is_none_or(|(s, _, _)| seq > *s) {
newest = Some((seq, model, path));
}
}
let Some((_, model, path)) = newest else {
return Ok(None);
};
let bytes = std::fs::read(&path).map_err(Error::Io)?;
Ok(report(&bytes, &model))
}
pub(super) fn model_entry(path: &Path) -> Option<(u32, String)> {
if path.extension().and_then(|e| e.to_str()) != Some("json") {
return None;
}
let stem = path.file_stem()?.to_string_lossy().into_owned();
let (seq, origin) = stem.split_once('-')?;
if origin == TOOL_ORIGIN || origin.is_empty() {
return None;
}
Some((seq.parse::<u32>().ok()?, origin.to_string()))
}
pub(super) fn report(bytes: &[u8], model: &str) -> Option<LastUsage> {
let usage = serde_json::from_slice::<Value>(bytes)
.ok()?
.get("usage")?
.clone();
let counter = |name: &str| usage.get(name).and_then(Value::as_u64).unwrap_or(0);
Some(LastUsage {
prompt_tokens: counter("input_tokens")
+ counter("cache_read_tokens")
+ counter("cache_write_tokens"),
context_window: usage.get("context_window").and_then(Value::as_u64),
model: model.to_string(),
})
}
#[cfg(test)]
mod tests;