use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::{LazyLock, Mutex};
use anyhow::{Context, Result};
use chrono::{DateTime, Utc};
use rusqlite::{Connection, OpenFlags};
use serde::Deserialize;
use serde_json::Value;
use crate::features::usage::{Row, Tool};
use crate::utils::files::num_f;
use crate::utils::time::parse_ts;
pub fn db_candidates() -> Vec<PathBuf> {
let over = std::env::var_os("TOKENBURN_ZED_DB")
.filter(|v| !v.is_empty())
.map(PathBuf::from);
let mut out: Vec<PathBuf> = over.into_iter().collect();
out.extend(candidates_from(
dirs::home_dir().as_deref(),
dirs::data_local_dir().as_deref(),
));
out
}
pub(crate) fn candidates_from(home: Option<&Path>, data_local: Option<&Path>) -> Vec<PathBuf> {
let mut out = Vec::new();
let mut push = |p: PathBuf| {
if !out.contains(&p) {
out.push(p);
}
};
if let Some(d) = data_local {
push(d.join("Zed/threads/threads.db"));
push(d.join("zed/threads/threads.db"));
}
if let Some(h) = home {
push(h.join("Library/Application Support/Zed/threads/threads.db"));
push(h.join(".local/share/zed/threads/threads.db"));
push(h.join(".var/app/dev.zed.Zed/data/zed/threads/threads.db"));
push(h.join(".config/zed/threads/threads.db"));
}
out
}
pub fn collect_zed(start: DateTime<Utc>) -> Result<Vec<Row>> {
match db_candidates().into_iter().find(|p| p.is_file()) {
Some(db) => collect_zed_from(&db, start),
None => Ok(Vec::new()),
}
}
type ThreadKey = (PathBuf, String);
type ThreadEntry = (String, Option<Row>);
type ThreadMap = HashMap<ThreadKey, ThreadEntry>;
static THREADS: LazyLock<Mutex<ThreadMap>> = LazyLock::new(|| Mutex::new(HashMap::new()));
fn threads() -> std::sync::MutexGuard<'static, ThreadMap> {
THREADS.lock().unwrap_or_else(|e| e.into_inner())
}
pub fn collect_zed_from(db: &Path, start: DateTime<Utc>) -> Result<Vec<Row>> {
let conn = Connection::open_with_flags(
db,
OpenFlags::SQLITE_OPEN_READ_ONLY | OpenFlags::SQLITE_OPEN_NO_MUTEX,
)
.with_context(|| format!("opening {}", db.display()))?;
let start_iso = start.format("%Y-%m-%dT%H:%M:%S").to_string();
let mut stmt = conn
.prepare("SELECT id, created_at, updated_at, data_type FROM threads WHERE created_at >= ?1")
.context("preparing threads query")?;
let metas: Vec<(String, String, String, Option<String>)> = stmt
.query_map([start_iso], |r| {
Ok((
r.get::<_, String>(0)?,
r.get::<_, String>(1)?,
r.get::<_, Option<String>>(2)?.unwrap_or_default(),
r.get::<_, Option<String>>(3)?,
))
})?
.flatten()
.collect();
let mut blob = conn
.prepare("SELECT data FROM threads WHERE id = ?1")
.context("preparing thread blob query")?;
let mut rows = Vec::new();
let mut live: Vec<ThreadKey> = Vec::with_capacity(metas.len());
for (id, created_at, updated_at, data_type) in metas {
let key = (db.to_path_buf(), id.clone());
live.push(key.clone());
let cached = threads()
.get(&key)
.filter(|(u, _)| *u == updated_at)
.map(|(_, row)| row.clone());
let row = match cached {
Some(row) => row,
None => {
let decoded = parse_ts(&created_at).and_then(|ts| {
let data: Vec<u8> = blob.query_row([&id], |r| r.get(0)).ok()?;
decode_thread(id, ts, data_type.as_deref(), &data)
});
threads().insert(key, (updated_at, decoded.clone()));
decoded
}
};
rows.extend(row);
}
let live: std::collections::HashSet<ThreadKey> = live.into_iter().collect();
threads().retain(|k, _| k.0 != db || live.contains(k));
Ok(rows)
}
#[cfg(test)]
pub(crate) fn cached_threads(db: &Path) -> usize {
threads().keys().filter(|k| k.0 == db).count()
}
#[derive(Deserialize)]
struct Thread {
model: Option<Value>,
cumulative_token_usage: Option<ThreadUsage>,
}
#[derive(Deserialize)]
struct ThreadUsage {
input_tokens: Option<f64>,
output_tokens: Option<f64>,
}
pub(crate) fn decode_thread(
id: String,
ts: DateTime<Utc>,
data_type: Option<&str>,
data: &[u8],
) -> Option<Row> {
let thread: Thread = if data_type == Some("zstd") {
serde_json::from_slice(&zstd::decode_all(data).ok()?).ok()?
} else {
serde_json::from_slice(data).ok()?
};
let usage = thread.cumulative_token_usage;
let model = thread
.model
.as_ref()
.and_then(|m| m.get("model"))
.and_then(Value::as_str)
.unwrap_or("unknown")
.to_string();
Some(Row {
tool: Tool::Zed,
project: model,
id,
ts,
input: num_f(usage.as_ref().and_then(|u| u.input_tokens)),
output: num_f(usage.as_ref().and_then(|u| u.output_tokens)),
cache_read: 0,
cache_write: 0,
cost: 0.0,
})
}
#[cfg(test)]
mod tests {
use super::*;
const JSON: &str = r#"{"model":{"model":"claude"},"cumulative_token_usage":{"input_tokens":7,"output_tokens":3}}"#;
#[test]
fn candidates_cover_every_platform_without_duplicates() {
let home = Path::new("/home/u");
let c = candidates_from(Some(home), Some(Path::new("/home/u/.local/share")));
let s: Vec<String> = c
.iter()
.map(|p| p.to_string_lossy().replace('\\', "/"))
.collect();
assert!(s.contains(&"/home/u/.local/share/zed/threads/threads.db".to_string()));
assert!(
s.contains(&"/home/u/Library/Application Support/Zed/threads/threads.db".to_string())
);
assert!(s.contains(&"/home/u/.var/app/dev.zed.Zed/data/zed/threads/threads.db".to_string()));
let n = s
.iter()
.filter(|p| p.ends_with(".local/share/zed/threads/threads.db"))
.count();
assert_eq!(n, 1);
}
#[test]
fn windows_style_local_appdata_is_used() {
let c = candidates_from(None, Some(Path::new("C:/Users/u/AppData/Local")));
assert_eq!(
c[0],
Path::new("C:/Users/u/AppData/Local").join("Zed/threads/threads.db")
);
}
#[test]
fn decodes_plain_and_zstd() {
let ts = "2026-05-01T10:00:00Z".parse().unwrap();
let plain = decode_thread("a".into(), ts, Some("json"), JSON.as_bytes()).unwrap();
assert_eq!((plain.input, plain.output), (7, 3));
assert_eq!(plain.project, "claude");
let packed = zstd::encode_all(JSON.as_bytes(), 0).unwrap();
let z = decode_thread("b".into(), ts, Some("zstd"), &packed).unwrap();
assert_eq!((z.input, z.output), (7, 3));
}
#[test]
fn only_changed_threads_are_decoded_again() {
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("threads.db");
let conn = Connection::open(&db).unwrap();
conn.execute(
"CREATE TABLE threads (id TEXT PRIMARY KEY, created_at TEXT, updated_at TEXT, data_type TEXT, data BLOB)",
[],
)
.unwrap();
let ins = |id: &str, updated: &str, json: &str| {
conn.execute(
"INSERT OR REPLACE INTO threads VALUES (?1, '2026-05-01T10:00:00Z', ?2, 'json', ?3)",
rusqlite::params![id, updated, json.as_bytes()],
)
.unwrap();
};
ins("t1", "u1", JSON);
ins("t2", "u1", JSON);
let start = "2026-01-01T00:00:00Z".parse().unwrap();
assert_eq!(collect_zed_from(&db, start).unwrap().len(), 2);
assert_eq!(cached_threads(&db), 2);
ins("t1", "u1", "not json");
ins(
"t2",
"u2",
r#"{"cumulative_token_usage":{"input_tokens":100,"output_tokens":1}}"#,
);
let rows = collect_zed_from(&db, start).unwrap();
assert_eq!(rows.len(), 2);
let t1 = rows.iter().find(|r| r.id == "t1").unwrap();
let t2 = rows.iter().find(|r| r.id == "t2").unwrap();
assert_eq!(t1.input, 7, "t1 came from the cache");
assert_eq!(
t2.input, 100,
"t2 was decoded again after its updated_at changed"
);
conn.execute("DELETE FROM threads WHERE id = 't1'", [])
.unwrap();
assert_eq!(collect_zed_from(&db, start).unwrap().len(), 1);
assert_eq!(cached_threads(&db), 1);
}
#[test]
fn reads_sqlite() {
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("threads.db");
let conn = Connection::open(&db).unwrap();
conn.execute(
"CREATE TABLE threads (id TEXT, created_at TEXT, updated_at TEXT, data_type TEXT, data BLOB)",
[],
)
.unwrap();
let packed = zstd::encode_all(JSON.as_bytes(), 0).unwrap();
conn.execute(
"INSERT INTO threads VALUES ('t1', '2026-05-01T10:00:00Z', 'u', 'zstd', ?1)",
[packed],
)
.unwrap();
drop(conn);
let start = "2026-01-01T00:00:00Z".parse().unwrap();
let rows = collect_zed_from(&db, start).unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].id, "t1");
}
}