Skip to main content

aft/db/
compression_events.rs

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/// Process-local compression totals backed by the durable event table.
90///
91/// Status reads validate entries with the table's maximum row id, an indexed
92/// lookup that detects writes from other AFT processes. Full aggregate scans run
93/// only for a cold or stale key. Successful local inserts advance warm entries
94/// directly while the caller still owns the database connection mutex.
95#[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    /// Apply a row that was inserted successfully on `conn`.
153    ///
154    /// A warm entry is advanced only when its watermark matches the row that
155    /// immediately preceded `inserted_row_id`. If another process wrote first,
156    /// the entry remains stale and the next status read rebuilds it from SQL.
157    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
217/// Insert one event and return its row id. Duplicate identities are ignored and
218/// return `None`, allowing in-process aggregates to advance only for durable rows.
219pub 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}