1use std::collections::HashMap;
2#[cfg(test)]
3use std::sync::atomic::{AtomicUsize, Ordering};
4
5use parking_lot::Mutex;
6use rusqlite::{params, Connection};
7
8pub struct CompressionEventRow<'a> {
9 pub harness: &'a str,
10 pub session_id: Option<&'a str>,
11 pub project_key: &'a str,
12 pub tool: &'a str,
13 pub task_id: Option<&'a str>,
14 pub command: Option<&'a str>,
15 pub compressor: &'a str,
16 pub original_bytes: i64,
17 pub compressed_bytes: i64,
18 pub original_tokens: u32,
19 pub compressed_tokens: u32,
20 pub created_at: i64,
21}
22
23#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize)]
24pub struct CompressionAggregate {
25 pub events: u64,
26 pub original_tokens: u64,
27 pub compressed_tokens: u64,
28}
29
30impl CompressionAggregate {
31 pub fn savings_tokens(&self) -> u64 {
32 self.original_tokens.saturating_sub(self.compressed_tokens)
33 }
34
35 fn add_event(&mut self, row: &CompressionEventRow<'_>) {
36 self.events = self.events.saturating_add(1);
37 self.original_tokens = self
38 .original_tokens
39 .saturating_add(u64::from(row.original_tokens));
40 self.compressed_tokens = self
41 .compressed_tokens
42 .saturating_add(u64::from(row.compressed_tokens));
43 }
44}
45
46#[derive(Debug, Clone, PartialEq, Eq, Hash)]
47struct ProjectAggregateKey {
48 harness: String,
49 project_key: String,
50}
51
52impl ProjectAggregateKey {
53 fn new(harness: &str, project_key: &str) -> Self {
54 Self {
55 harness: harness.to_string(),
56 project_key: project_key.to_string(),
57 }
58 }
59}
60
61#[derive(Debug, Clone, PartialEq, Eq, Hash)]
62struct SessionAggregateKey {
63 project: ProjectAggregateKey,
64 session_id: String,
65}
66
67impl SessionAggregateKey {
68 fn new(harness: &str, project_key: &str, session_id: &str) -> Self {
69 Self {
70 project: ProjectAggregateKey::new(harness, project_key),
71 session_id: session_id.to_string(),
72 }
73 }
74}
75
76#[derive(Debug, Clone, Copy)]
77struct CachedAggregate {
78 aggregate: CompressionAggregate,
79 watermark: i64,
80}
81
82#[derive(Debug, Default)]
83struct CompressionAggregateCacheInner {
84 connection_identity: Option<usize>,
85 projects: HashMap<ProjectAggregateKey, CachedAggregate>,
86 sessions: HashMap<SessionAggregateKey, CachedAggregate>,
87}
88
89#[derive(Debug, Default)]
96pub struct CompressionAggregateCache {
97 inner: Mutex<CompressionAggregateCacheInner>,
98 #[cfg(test)]
99 aggregate_scan_count: AtomicUsize,
100}
101
102impl CompressionAggregateCache {
103 pub fn aggregates_for_session(
104 &self,
105 conn: &Connection,
106 harness: &str,
107 project_key: &str,
108 session_id: &str,
109 ) -> rusqlite::Result<(CompressionAggregate, CompressionAggregate)> {
110 let watermark = compression_event_watermark(conn)?;
111 let project_key = ProjectAggregateKey::new(harness, project_key);
112 let session_key = SessionAggregateKey::new(harness, &project_key.project_key, session_id);
113 let mut inner = self.inner.lock();
114 reset_for_connection_change(&mut inner, conn);
115
116 let project = match inner.projects.get(&project_key) {
117 Some(cached) if cached.watermark == watermark => cached.aggregate,
118 _ => {
119 self.note_aggregate_scan();
120 let aggregate = aggregate_for_project(conn, harness, &project_key.project_key)?;
121 inner.projects.insert(
122 project_key.clone(),
123 CachedAggregate {
124 aggregate,
125 watermark,
126 },
127 );
128 aggregate
129 }
130 };
131
132 let session = match inner.sessions.get(&session_key) {
133 Some(cached) if cached.watermark == watermark => cached.aggregate,
134 _ => {
135 self.note_aggregate_scan();
136 let aggregate =
137 aggregate_for_session(conn, harness, &project_key.project_key, session_id)?;
138 inner.sessions.insert(
139 session_key,
140 CachedAggregate {
141 aggregate,
142 watermark,
143 },
144 );
145 aggregate
146 }
147 };
148
149 Ok((project, session))
150 }
151
152 pub fn record_successful_insert(
158 &self,
159 conn: &Connection,
160 row: &CompressionEventRow<'_>,
161 inserted_row_id: i64,
162 ) {
163 let previous_watermark = compression_event_watermark_before(conn, inserted_row_id);
164 let project_key = ProjectAggregateKey::new(row.harness, row.project_key);
165 let session_key = row
166 .session_id
167 .map(|session_id| SessionAggregateKey::new(row.harness, row.project_key, session_id));
168 let mut inner = self.inner.lock();
169 reset_for_connection_change(&mut inner, conn);
170
171 let Ok(previous_watermark) = previous_watermark else {
172 *inner = CompressionAggregateCacheInner {
173 connection_identity: inner.connection_identity,
174 ..CompressionAggregateCacheInner::default()
175 };
176 return;
177 };
178
179 for (key, cached) in &mut inner.projects {
180 if cached.watermark != previous_watermark {
181 continue;
182 }
183 if key == &project_key {
184 cached.aggregate.add_event(row);
185 }
186 cached.watermark = inserted_row_id;
187 }
188 for (key, cached) in &mut inner.sessions {
189 if cached.watermark != previous_watermark {
190 continue;
191 }
192 if session_key.as_ref() == Some(key) {
193 cached.aggregate.add_event(row);
194 }
195 cached.watermark = inserted_row_id;
196 }
197 }
198
199 pub fn clear(&self) {
200 *self.inner.lock() = CompressionAggregateCacheInner::default();
201 }
202
203 #[cfg(test)]
204 fn aggregate_scan_count_for_test(&self) -> usize {
205 self.aggregate_scan_count.load(Ordering::Relaxed)
206 }
207
208 #[cfg(test)]
209 fn note_aggregate_scan(&self) {
210 self.aggregate_scan_count.fetch_add(1, Ordering::Relaxed);
211 }
212
213 #[cfg(not(test))]
214 fn note_aggregate_scan(&self) {}
215}
216
217pub fn insert_compression_event(
220 conn: &Connection,
221 row: &CompressionEventRow<'_>,
222) -> rusqlite::Result<Option<i64>> {
223 let inserted = conn.execute(
224 r#"
225 INSERT OR IGNORE INTO compression_events (
226 harness, session_id, project_key, tool, task_id, command, compressor,
227 original_bytes, compressed_bytes, original_tokens, compressed_tokens, created_at
228 )
229 VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12)
230 "#,
231 params![
232 row.harness,
233 row.session_id,
234 row.project_key,
235 row.tool,
236 row.task_id,
237 row.command,
238 row.compressor,
239 row.original_bytes,
240 row.compressed_bytes,
241 row.original_tokens,
242 row.compressed_tokens,
243 row.created_at,
244 ],
245 )?;
246 Ok((inserted > 0).then(|| conn.last_insert_rowid()))
247}
248
249pub fn aggregate_for_project(
250 conn: &Connection,
251 harness: &str,
252 project_key: &str,
253) -> rusqlite::Result<CompressionAggregate> {
254 conn.query_row(
255 r#"
256 SELECT
257 COUNT(*) AS events,
258 COALESCE(SUM(original_tokens), 0) AS original,
259 COALESCE(SUM(compressed_tokens), 0) AS compressed
260 FROM compression_events
261 WHERE harness = ?1 AND project_key = ?2
262 "#,
263 params![harness, project_key],
264 |row| {
265 Ok(CompressionAggregate {
266 events: row.get::<_, i64>(0)? as u64,
267 original_tokens: row.get::<_, i64>(1)? as u64,
268 compressed_tokens: row.get::<_, i64>(2)? as u64,
269 })
270 },
271 )
272}
273
274pub fn aggregate_for_session(
275 conn: &Connection,
276 harness: &str,
277 project_key: &str,
278 session_id: &str,
279) -> rusqlite::Result<CompressionAggregate> {
280 conn.query_row(
281 r#"
282 SELECT
283 COUNT(*) AS events,
284 COALESCE(SUM(original_tokens), 0) AS original,
285 COALESCE(SUM(compressed_tokens), 0) AS compressed
286 FROM compression_events
287 WHERE harness = ?1 AND project_key = ?2 AND session_id = ?3
288 "#,
289 params![harness, project_key, session_id],
290 |row| {
291 Ok(CompressionAggregate {
292 events: row.get::<_, i64>(0)? as u64,
293 original_tokens: row.get::<_, i64>(1)? as u64,
294 compressed_tokens: row.get::<_, i64>(2)? as u64,
295 })
296 },
297 )
298}
299
300fn reset_for_connection_change(inner: &mut CompressionAggregateCacheInner, conn: &Connection) {
301 let identity = conn as *const Connection as usize;
302 if inner.connection_identity != Some(identity) {
303 *inner = CompressionAggregateCacheInner {
304 connection_identity: Some(identity),
305 ..CompressionAggregateCacheInner::default()
306 };
307 }
308}
309
310fn compression_event_watermark(conn: &Connection) -> rusqlite::Result<i64> {
311 conn.query_row(
312 "SELECT COALESCE(MAX(id), 0) FROM compression_events",
313 [],
314 |row| row.get(0),
315 )
316}
317
318fn compression_event_watermark_before(
319 conn: &Connection,
320 inserted_row_id: i64,
321) -> rusqlite::Result<i64> {
322 conn.query_row(
323 "SELECT COALESCE(MAX(id), 0) FROM compression_events WHERE id < ?1",
324 [inserted_row_id],
325 |row| row.get(0),
326 )
327}
328
329#[cfg(test)]
330mod tests {
331 use super::*;
332 use tempfile::tempdir;
333
334 #[test]
335 fn duplicate_identity_is_ignored_without_cross_project_suppression() {
336 let dir = tempdir().expect("tempdir");
337 let conn = crate::db::open(&dir.path().join("aft.db")).expect("open db");
338
339 assert!(
340 insert_compression_event(&conn, &row("project-a", "task-1", 100, 40, 1))
341 .expect("insert first")
342 .is_some()
343 );
344 assert!(
345 insert_compression_event(&conn, &row("project-a", "task-1", 900, 10, 2))
346 .expect("ignore duplicate")
347 .is_none()
348 );
349 assert!(
350 insert_compression_event(&conn, &row("project-b", "task-1", 200, 80, 3))
351 .expect("insert same task id for other project")
352 .is_some()
353 );
354
355 let project_a = aggregate_for_project(&conn, "opencode", "project-a").unwrap();
356 assert_eq!(project_a.events, 1);
357 assert_eq!(project_a.original_tokens, 100);
358 assert_eq!(project_a.compressed_tokens, 40);
359
360 let project_b = aggregate_for_project(&conn, "opencode", "project-b").unwrap();
361 assert_eq!(project_b.events, 1);
362 assert_eq!(project_b.original_tokens, 200);
363 assert_eq!(project_b.compressed_tokens, 80);
364 }
365
366 #[test]
367 fn cached_aggregates_match_sql_after_generated_inserts_and_duplicates() {
368 let dir = tempdir().expect("tempdir");
369 let conn = crate::db::open(&dir.path().join("aft.db")).expect("open db");
370 let cache = CompressionAggregateCache::default();
371 let (project, session) = cache
372 .aggregates_for_session(&conn, "opencode", "project-a", "session-1")
373 .expect("warm cache");
374 assert_eq!(project, CompressionAggregate::default());
375 assert_eq!(session, CompressionAggregate::default());
376 cache
377 .aggregates_for_session(&conn, "opencode", "project-a", "session-2")
378 .expect("warm sibling session");
379 cache
380 .aggregates_for_session(&conn, "opencode", "project-b", "session-1")
381 .expect("warm sibling project");
382 assert_eq!(cache.aggregate_scan_count_for_test(), 5);
383
384 let mut previous_task = String::new();
385 for index in 0..64u32 {
386 let task_id = if index % 5 == 4 {
387 previous_task.clone()
388 } else {
389 let task_id = format!("task-{index}");
390 previous_task = task_id.clone();
391 task_id
392 };
393 let row = row(
394 "project-a",
395 &task_id,
396 100 + index,
397 40 + (index % 17),
398 i64::from(index),
399 );
400 if let Some(row_id) = insert_compression_event(&conn, &row).expect("insert event") {
401 cache.record_successful_insert(&conn, &row, row_id);
402 }
403
404 for (project_key, session_id) in [
405 ("project-a", "session-1"),
406 ("project-a", "session-2"),
407 ("project-b", "session-1"),
408 ] {
409 let cached = cache
410 .aggregates_for_session(&conn, "opencode", project_key, session_id)
411 .expect("read cache");
412 let scanned = (
413 aggregate_for_project(&conn, "opencode", project_key).expect("scan project"),
414 aggregate_for_session(&conn, "opencode", project_key, session_id)
415 .expect("scan session"),
416 );
417 assert_eq!(cached, scanned, "aggregate mismatch after step {index}");
418 }
419 assert_eq!(
420 cache.aggregate_scan_count_for_test(),
421 5,
422 "local inserts must advance warm entries without rescanning"
423 );
424 }
425 }
426
427 fn row<'a>(
428 project_key: &'a str,
429 task_id: &'a str,
430 original_tokens: u32,
431 compressed_tokens: u32,
432 created_at: i64,
433 ) -> CompressionEventRow<'a> {
434 CompressionEventRow {
435 harness: "opencode",
436 session_id: Some("session-1"),
437 project_key,
438 tool: "bash",
439 task_id: Some(task_id),
440 command: Some("echo ok"),
441 compressor: "zstd",
442 original_bytes: i64::from(original_tokens) * 4,
443 compressed_bytes: i64::from(compressed_tokens) * 4,
444 original_tokens,
445 compressed_tokens,
446 created_at,
447 }
448 }
449}