use anyhow::{Context, Result};
use std::fs::OpenOptions;
use std::io::{BufWriter, Write};
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use tokio::sync::{mpsc, oneshot};
use tokio::task::spawn_blocking;
use tokio::time::{self, MissedTickBehavior};
use crate::config::constants::defaults;
use crate::utils::file_utils::ensure_dir_exists_sync;
enum LogMessage {
Line { line: String, bytes: usize },
Flush(oneshot::Sender<bool>),
}
#[derive(Default)]
struct Diagnostics {
dropped_lines: AtomicUsize,
dropped_bytes: AtomicUsize,
write_failures: AtomicUsize,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[allow(dead_code)]
pub(crate) struct AsyncLineWriterDiagnostics {
pub dropped_lines: usize,
pub dropped_bytes: usize,
pub write_failures: usize,
}
#[derive(Clone)]
pub struct AsyncLineWriter {
sender: mpsc::Sender<LogMessage>,
reserved_bytes: Arc<AtomicUsize>,
diagnostics: Arc<Diagnostics>,
max_buffer_bytes: usize,
}
impl AsyncLineWriter {
pub fn new(path: PathBuf) -> Result<Self> {
Self::new_with_limits(
path,
defaults::DEFAULT_TRAJECTORY_LOG_CHANNEL_CAPACITY,
defaults::DEFAULT_TRAJECTORY_LOG_BUFFER_BYTES,
Duration::from_millis(defaults::DEFAULT_TRAJECTORY_LOG_FLUSH_INTERVAL_MS),
)
}
fn new_with_limits(
path: PathBuf,
channel_capacity: usize,
buffer_bytes: usize,
flush_interval: Duration,
) -> Result<Self> {
if let Some(parent) = path.parent() {
ensure_dir_exists_sync(parent)
.with_context(|| format!("Failed to create log directory: {}", parent.display()))?;
}
OpenOptions::new()
.create(true)
.append(true)
.open(&path)
.with_context(|| format!("Failed to create log file: {}", path.display()))?;
let (sender, receiver) = mpsc::channel(channel_capacity.max(1));
let reserved_bytes = Arc::new(AtomicUsize::new(0));
let diagnostics = Arc::new(Diagnostics::default());
let actor_reserved_bytes = Arc::clone(&reserved_bytes);
let actor_diagnostics = Arc::clone(&diagnostics);
tokio::spawn(async move {
actor_task(
&path,
receiver,
actor_reserved_bytes,
actor_diagnostics,
channel_capacity.max(1),
buffer_bytes,
flush_interval,
)
.await;
});
Ok(Self {
sender,
reserved_bytes,
diagnostics,
max_buffer_bytes: buffer_bytes.max(1),
})
}
pub fn write_line(&self, line: String) {
let bytes = line.len().saturating_add(1);
if !reserve_bytes(&self.reserved_bytes, bytes, self.max_buffer_bytes) {
self.record_drop(bytes);
return;
}
if self.sender.try_send(LogMessage::Line { line, bytes }).is_err() {
self.reserved_bytes.fetch_sub(bytes, Ordering::AcqRel);
self.record_drop(bytes);
}
}
pub async fn flush(&self) {
let _ = self.flush_result().await;
}
pub(crate) async fn flush_result(&self) -> std::io::Result<()> {
let (tx, rx) = oneshot::channel();
self.sender
.send(LogMessage::Flush(tx))
.await
.map_err(|_error| std::io::Error::other("trajectory writer actor closed"))?;
match rx
.await
.map_err(|_error| std::io::Error::other("trajectory writer actor dropped flush acknowledgement"))?
{
true => Ok(()),
false => Err(std::io::Error::other("trajectory writer failed to flush a batch")),
}
}
#[allow(dead_code)]
pub(crate) fn diagnostics(&self) -> AsyncLineWriterDiagnostics {
AsyncLineWriterDiagnostics {
dropped_lines: self.diagnostics.dropped_lines.load(Ordering::Relaxed),
dropped_bytes: self.diagnostics.dropped_bytes.load(Ordering::Relaxed),
write_failures: self.diagnostics.write_failures.load(Ordering::Relaxed),
}
}
fn record_drop(&self, bytes: usize) {
self.diagnostics.dropped_lines.fetch_add(1, Ordering::Relaxed);
self.diagnostics.dropped_bytes.fetch_add(bytes, Ordering::Relaxed);
tracing::debug!(bytes, "dropping trajectory line because the writer is saturated");
}
}
fn reserve_bytes(reserved_bytes: &AtomicUsize, bytes: usize, max_bytes: usize) -> bool {
reserved_bytes
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| {
current.checked_add(bytes).filter(|next| *next <= max_bytes)
})
.is_ok()
}
async fn actor_task(
path: &Path,
mut receiver: mpsc::Receiver<LogMessage>,
reserved_bytes: Arc<AtomicUsize>,
diagnostics: Arc<Diagnostics>,
buffer_lines: usize,
buffer_bytes: usize,
flush_interval: Duration,
) {
let mut buffer: Vec<(String, usize)> = Vec::new();
let mut buffered_bytes = 0usize;
let mut interval = time::interval(flush_interval.max(Duration::from_millis(1)));
interval.set_missed_tick_behavior(MissedTickBehavior::Delay);
loop {
tokio::select! {
msg = receiver.recv() => match msg {
Some(LogMessage::Line { line, bytes }) => {
buffered_bytes = buffered_bytes.saturating_add(bytes);
buffer.push((line, bytes));
if buffer.len() >= buffer_lines
|| buffered_bytes >= buffer_bytes
{
let _ = flush_buffer(path, &mut buffer, &mut buffered_bytes, &reserved_bytes, &diagnostics).await;
}
}
Some(LogMessage::Flush(ack)) => {
let flushed = flush_buffer(path, &mut buffer, &mut buffered_bytes, &reserved_bytes, &diagnostics).await;
let _ = ack.send(flushed);
}
None => break,
},
_ = interval.tick() => {
let _ = flush_buffer(path, &mut buffer, &mut buffered_bytes, &reserved_bytes, &diagnostics).await;
},
}
}
let _ = flush_buffer(path, &mut buffer, &mut buffered_bytes, &reserved_bytes, &diagnostics).await;
}
async fn flush_buffer(
path: &Path,
buffer: &mut Vec<(String, usize)>,
buffered_bytes: &mut usize,
reserved_bytes: &AtomicUsize,
diagnostics: &Diagnostics,
) -> bool {
if buffer.is_empty() {
return true;
}
let lines = std::mem::take(buffer);
let line_count = lines.len();
let bytes = *buffered_bytes;
*buffered_bytes = 0;
let written = flush_lines(path, lines).await;
reserved_bytes.fetch_sub(bytes, Ordering::AcqRel);
if !written {
diagnostics.dropped_lines.fetch_add(line_count, Ordering::Relaxed);
diagnostics.dropped_bytes.fetch_add(bytes, Ordering::Relaxed);
diagnostics.write_failures.fetch_add(1, Ordering::Relaxed);
tracing::warn!(line_count, bytes, "trajectory writer failed to flush a batch");
}
written
}
async fn flush_lines(path: &Path, lines: Vec<(String, usize)>) -> bool {
let path = path.to_path_buf();
spawn_blocking(move || {
let file = match OpenOptions::new().create(true).append(true).open(&path) {
Ok(file) => file,
Err(_) => return false,
};
let mut writer = BufWriter::new(file);
for (line, _) in &lines {
if writeln!(writer, "{line}").is_err() {
return false;
}
}
writer.flush().is_ok()
})
.await
.unwrap_or(false)
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
use tempfile::TempDir;
#[tokio::test]
async fn drops_lines_when_byte_bound_is_saturated() {
let temp_dir = TempDir::new().expect("temp directory");
let writer =
AsyncLineWriter::new_with_limits(temp_dir.path().join("trajectory.jsonl"), 4, 4, Duration::from_secs(60))
.expect("writer");
writer.write_line("1234".to_owned());
writer.write_line("more data".to_owned());
let diagnostics = writer.diagnostics();
assert_eq!(diagnostics.dropped_lines, 2);
assert_eq!(diagnostics.dropped_bytes, 15);
}
#[tokio::test]
async fn periodic_flush_writes_without_explicit_flush() {
let temp_dir = TempDir::new().expect("temp directory");
let path = temp_dir.path().join("trajectory.jsonl");
let writer =
AsyncLineWriter::new_with_limits(path.clone(), 4, 1024, Duration::from_millis(10)).expect("writer");
writer.write_line("periodic".to_owned());
time::sleep(Duration::from_millis(40)).await;
assert_eq!(fs::read_to_string(path).expect("read log"), "periodic\n");
}
#[tokio::test]
async fn channel_shutdown_flushes_pending_lines() {
let temp_dir = TempDir::new().expect("temp directory");
let path = temp_dir.path().join("trajectory.jsonl");
{
let writer =
AsyncLineWriter::new_with_limits(path.clone(), 4, 1024, Duration::from_secs(60)).expect("writer");
writer.write_line("final".to_owned());
}
for _ in 0..20 {
if fs::read_to_string(&path).expect("read log") == "final\n" {
return;
}
time::sleep(Duration::from_millis(5)).await;
}
panic!("actor did not flush before shutdown");
}
#[tokio::test]
async fn flush_result_reports_file_write_failure() {
let temp_dir = TempDir::new().expect("temp directory");
let path = temp_dir.path().join("trajectory.jsonl");
let writer = AsyncLineWriter::new_with_limits(path, 4, 1024, Duration::from_secs(60)).expect("writer");
fs::remove_dir_all(temp_dir.path()).expect("remove temporary directory");
writer.write_line("unwritable".to_owned());
assert!(writer.flush_result().await.is_err());
let diagnostics = writer.diagnostics();
assert_eq!(diagnostics.write_failures, 1);
assert_eq!(diagnostics.dropped_lines, 1);
assert_eq!(diagnostics.dropped_bytes, "unwritable\n".len());
}
}