diff --git a/codex-rs/rollout/src/lib.rs b/codex-rs/rollout/src/lib.rs index 3c103750c178..43e7a9216f05 100644 --- a/codex-rs/rollout/src/lib.rs +++ b/codex-rs/rollout/src/lib.rs @@ -11,6 +11,7 @@ pub(crate) mod metadata; mod persistence_metrics; pub(crate) mod policy; pub(crate) mod recorder; +mod reverse_jsonl_scanner; pub(crate) mod search; pub(crate) mod session_index; mod sqlite_metrics; diff --git a/codex-rs/rollout/src/reverse_jsonl_scanner.rs b/codex-rs/rollout/src/reverse_jsonl_scanner.rs new file mode 100644 index 000000000000..c84809ca1ae2 --- /dev/null +++ b/codex-rs/rollout/src/reverse_jsonl_scanner.rs @@ -0,0 +1,103 @@ +use std::io; +use std::io::Read; +use std::io::Seek; +use std::io::SeekFrom; + +use serde::de::DeserializeOwned; + +const READ_CHUNK_SIZE: usize = 8 * 1024; + +#[derive(Debug)] +pub(crate) enum ScanOutcome { + Parsed(T), + #[allow(dead_code)] + Rejected(serde_json::Error), +} + +/// Read-only scanner for newline-delimited JSON records, starting from the end. +pub(crate) struct ReverseJsonlScanner { + reader: R, + next_chunk_end: u64, + chunk_position: usize, + chunk: Vec, + record_reversed: Vec, +} + +impl ReverseJsonlScanner +where + R: Read + Seek, +{ + pub(crate) fn new(mut reader: R) -> io::Result { + let next_chunk_end = reader.seek(SeekFrom::End(0))?; + Ok(Self { + reader, + next_chunk_end, + chunk_position: 0, + chunk: vec![0; READ_CHUNK_SIZE], + record_reversed: Vec::new(), + }) + } + + /// Scans the next nonblank record. + /// + /// I/O failures are returned as [`Err`]. Invalid JSON records are returned as + /// [`ScanOutcome::Rejected`], and the scanner remains usable. + pub(crate) fn scan_next(&mut self) -> io::Result>> + where + T: DeserializeOwned, + { + loop { + let Some(byte) = self.read_previous_byte()? else { + return Ok(self.finish_record()); + }; + + if byte != b'\n' { + self.record_reversed.push(byte); + continue; + } + + if let Some(outcome) = self.finish_record() { + return Ok(Some(outcome)); + } + } + } + + fn read_previous_byte(&mut self) -> io::Result> { + if self.chunk_position == 0 { + if self.next_chunk_end == 0 { + return Ok(None); + } + + let read_size = usize::try_from(self.next_chunk_end.min(READ_CHUNK_SIZE as u64)) + .map_err(io::Error::other)?; + self.next_chunk_end -= read_size as u64; + self.reader.seek(SeekFrom::Start(self.next_chunk_end))?; + self.reader.read_exact(&mut self.chunk[..read_size])?; + self.chunk_position = read_size; + } + + self.chunk_position -= 1; + Ok(Some(self.chunk[self.chunk_position])) + } + + fn finish_record(&mut self) -> Option> + where + T: DeserializeOwned, + { + self.record_reversed.reverse(); + let outcome = if self.record_reversed.iter().all(u8::is_ascii_whitespace) { + None + } else { + Some(match serde_json::from_slice::(&self.record_reversed) { + Ok(value) => ScanOutcome::Parsed(value), + Err(error) => ScanOutcome::Rejected(error), + }) + }; + self.record_reversed.clear(); + outcome + } +} + +#[cfg(test)] +#[path = "reverse_jsonl_scanner_tests.rs"] +mod tests; diff --git a/codex-rs/rollout/src/reverse_jsonl_scanner_tests.rs b/codex-rs/rollout/src/reverse_jsonl_scanner_tests.rs new file mode 100644 index 000000000000..0e83239397f9 --- /dev/null +++ b/codex-rs/rollout/src/reverse_jsonl_scanner_tests.rs @@ -0,0 +1,141 @@ +use std::io::Cursor; +use std::io::Read; +use std::io::Seek; + +use pretty_assertions::assert_eq; +use serde::Deserialize; +use serde::Serialize; + +use super::ReverseJsonlScanner; +use super::ScanOutcome; + +#[derive(Debug, Deserialize, Serialize, PartialEq)] +struct TestRecord { + value: String, +} + +fn record(value: &str) -> TestRecord { + TestRecord { + value: value.to_string(), + } +} + +fn parsed(outcome: Option>) -> T { + let Some(ScanOutcome::Parsed(record)) = outcome else { + panic!("expected parsed record"); + }; + record +} + +fn assert_records(scanner: &mut ReverseJsonlScanner, expected: &[&str]) -> std::io::Result<()> +where + R: Read + Seek, +{ + for value in expected { + assert_eq!(parsed(scanner.scan_next::()?), record(value)); + } + assert!(scanner.scan_next::()?.is_none()); + Ok(()) +} + +#[test] +fn scans_jsonl_records_from_end() -> std::io::Result<()> { + let input = br#"{"value":"first"} +{"value":"second"} +{"value":"third"} +"#; + + assert_records( + &mut ReverseJsonlScanner::new(Cursor::new(input))?, + &["third", "second", "first"], + ) +} + +#[test] +fn rejects_invalid_json_and_continues_scanning() -> std::io::Result<()> { + let input = br#"{"value":"first"} +not-json +{"value":"third"} +"#; + let mut scanner = ReverseJsonlScanner::new(Cursor::new(input))?; + + assert_eq!(parsed(scanner.scan_next::()?), record("third")); + let Some(ScanOutcome::Rejected(error)) = scanner.scan_next::()? else { + panic!("expected rejected record"); + }; + assert!(error.is_syntax()); + assert_eq!(parsed(scanner.scan_next::()?), record("first")); + Ok(()) +} + +#[test] +fn accepts_valid_json_at_eof() -> std::io::Result<()> { + let input = b"{\"value\":\"first\"}\n{\"value\":\"second\"}"; + + assert_records( + &mut ReverseJsonlScanner::new(Cursor::new(input))?, + &["second", "first"], + ) +} + +#[test] +fn rejects_invalid_json_at_eof_and_continues_scanning() -> std::io::Result<()> { + let input = b"{\"value\":\"first\"}\n{\"value\":"; + let mut scanner = ReverseJsonlScanner::new(Cursor::new(input))?; + + let Some(ScanOutcome::Rejected(error)) = scanner.scan_next::()? else { + panic!("expected rejected record"); + }; + assert!(error.is_eof()); + assert_eq!(parsed(scanner.scan_next::()?), record("first")); + Ok(()) +} + +#[test] +fn skips_blank_lines_with_or_without_termination() -> std::io::Result<()> { + let input = b"{\"value\":\"first\"}\r\n\n \t\r"; + + assert_records( + &mut ReverseJsonlScanner::new(Cursor::new(input))?, + &["first"], + ) +} + +#[test] +fn scans_across_read_chunk_boundaries() -> std::io::Result<()> { + let empty_record_len = serde_json::to_string(&record(""))?.len(); + for distance_from_eof in [ + super::READ_CHUNK_SIZE - 1, + super::READ_CHUNK_SIZE, + super::READ_CHUNK_SIZE + 1, + ] { + let large_value = "x".repeat(distance_from_eof - empty_record_len - 2); + let input = format!( + "{}\n{}\n", + serde_json::to_string(&record("first"))?, + serde_json::to_string(&record(&large_value))? + ); + let mut scanner = ReverseJsonlScanner::new(Cursor::new(input.into_bytes()))?; + + assert_eq!( + parsed(scanner.scan_next::()?), + record(&large_value) + ); + assert_eq!(parsed(scanner.scan_next::()?), record("first")); + } + Ok(()) +} + +#[test] +fn scans_record_spanning_three_read_chunks() -> std::io::Result<()> { + let large_value = "x".repeat(super::READ_CHUNK_SIZE * 2); + let input = format!( + "{}\n{}\n{}\n", + serde_json::to_string(&record("first"))?, + serde_json::to_string(&record(&large_value))?, + serde_json::to_string(&record("third"))? + ); + let mut scanner = ReverseJsonlScanner::new(Cursor::new(input.into_bytes()))?; + + assert_records(&mut scanner, &["third", &large_value, "first"]) +} diff --git a/codex-rs/rollout/src/session_index.rs b/codex-rs/rollout/src/session_index.rs index 96ff9d825e01..834dd8a70d08 100644 --- a/codex-rs/rollout/src/session_index.rs +++ b/codex-rs/rollout/src/session_index.rs @@ -2,15 +2,14 @@ use std::collections::HashMap; use std::collections::HashSet; use std::fs::File; use std::io::ErrorKind; -use std::io::Read; -use std::io::Seek; -use std::io::SeekFrom; use std::io::Write; use std::path::Path; use std::path::PathBuf; use std::sync::LazyLock; use std::sync::Mutex; +use crate::reverse_jsonl_scanner::ReverseJsonlScanner; +use crate::reverse_jsonl_scanner::ScanOutcome; use codex_protocol::ThreadId; use codex_protocol::protocol::SessionMetaLine; use serde::Deserialize; @@ -18,7 +17,6 @@ use serde::Serialize; use tokio::io::AsyncBufReadExt; const SESSION_INDEX_FILE: &str = "session_index.jsonl"; -const READ_CHUNK_SIZE: usize = 8192; static SESSION_INDEX_LOCK: LazyLock> = LazyLock::new(|| Mutex::new(())); #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] @@ -245,64 +243,18 @@ fn scan_index_from_end_for_each( where F: FnMut(&SessionIndexEntry) -> std::io::Result>, { - let mut file = File::open(path)?; - let mut remaining = file.metadata()?.len(); - let mut line_rev: Vec = Vec::new(); - let mut buf = vec![0u8; READ_CHUNK_SIZE]; - - while remaining > 0 { - let read_size = usize::try_from(remaining.min(READ_CHUNK_SIZE as u64)) - .map_err(std::io::Error::other)?; - remaining -= read_size as u64; - file.seek(SeekFrom::Start(remaining))?; - file.read_exact(&mut buf[..read_size])?; - - for &byte in buf[..read_size].iter().rev() { - if byte == b'\n' { - if let Some(entry) = parse_line_from_rev(&mut line_rev, &mut visit_entry)? { - return Ok(Some(entry)); - } - continue; - } - line_rev.push(byte); + let mut scanner = ReverseJsonlScanner::new(File::open(path)?)?; + while let Some(outcome) = scanner.scan_next::()? { + let ScanOutcome::Parsed(entry) = outcome else { + continue; + }; + if let Some(entry) = visit_entry(&entry)? { + return Ok(Some(entry)); } } - - if let Some(entry) = parse_line_from_rev(&mut line_rev, &mut visit_entry)? { - return Ok(Some(entry)); - } - Ok(None) } -fn parse_line_from_rev( - line_rev: &mut Vec, - visit_entry: &mut F, -) -> std::io::Result> -where - F: FnMut(&SessionIndexEntry) -> std::io::Result>, -{ - if line_rev.is_empty() { - return Ok(None); - } - line_rev.reverse(); - let line = std::mem::take(line_rev); - let Ok(mut line) = String::from_utf8(line) else { - return Ok(None); - }; - if line.ends_with('\r') { - line.pop(); - } - let trimmed = line.trim(); - if trimmed.is_empty() { - return Ok(None); - } - let Ok(entry) = serde_json::from_str::(trimmed) else { - return Ok(None); - }; - visit_entry(&entry) -} - #[cfg(test)] #[path = "session_index_tests.rs"] mod tests; diff --git a/codex-rs/rollout/src/session_index_tests.rs b/codex-rs/rollout/src/session_index_tests.rs index 6eb1fa037463..45aad8279cfb 100644 --- a/codex-rs/rollout/src/session_index_tests.rs +++ b/codex-rs/rollout/src/session_index_tests.rs @@ -234,6 +234,40 @@ fn scan_index_returns_none_when_entry_missing() -> std::io::Result<()> { Ok(()) } +#[tokio::test] +async fn reverse_lookup_accepts_valid_eof_json_and_skips_invalid() -> std::io::Result<()> { + let temp = TempDir::new()?; + let path = session_index_path(temp.path()); + let expected = SessionIndexEntry { + id: ThreadId::new(), + thread_name: "expected".to_string(), + updated_at: "2024-01-01T00:00:00Z".to_string(), + }; + let unterminated = SessionIndexEntry { + id: ThreadId::new(), + thread_name: "unterminated".to_string(), + updated_at: "2024-01-02T00:00:00Z".to_string(), + }; + std::fs::write( + &path, + format!( + "{}\nnot-json\n{}", + serde_json::to_string(&expected)?, + serde_json::to_string(&unterminated)? + ), + )?; + + assert_eq!( + find_thread_name_by_id(temp.path(), &unterminated.id).await?, + Some("unterminated".to_string()) + ); + assert_eq!( + find_thread_name_by_id(temp.path(), &expected.id).await?, + Some("expected".to_string()) + ); + Ok(()) +} + #[tokio::test] async fn find_thread_names_by_ids_prefers_latest_entry() -> std::io::Result<()> { let temp = TempDir::new()?;