use async_trait::async_trait;
use std::collections::VecDeque;
use std::io::Write as _;
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::mpsc::{SyncSender, TrySendError};
use std::thread::JoinHandle;
use tokio::fs::{File, OpenOptions};
use tokio::io::AsyncWriteExt;
use tokio::sync::Mutex;
use super::events::AuditEvent;
use super::types::{AuditError, AuditResult};
#[async_trait]
pub trait AuditOutput: Send + Sync {
fn name(&self) -> &str;
async fn write(&self, event: &AuditEvent) -> AuditResult<()>;
async fn flush(&self) -> AuditResult<()>;
async fn close(&self) -> AuditResult<()>;
}
pub type BoxedAuditOutput = Box<dyn AuditOutput>;
enum StderrCommand {
Write {
serialized: String,
completion: tokio::sync::oneshot::Sender<std::io::Result<()>>,
},
Close,
}
pub struct StderrOutput {
sender: SyncSender<StderrCommand>,
worker: Mutex<Option<JoinHandle<std::io::Result<()>>>>,
}
impl StderrOutput {
pub fn new(buffer_size: usize) -> AuditResult<Self> {
let (sender, receiver) = std::sync::mpsc::sync_channel(buffer_size);
let worker = std::thread::Builder::new()
.name("audit-stderr-writer".to_string())
.spawn(move || {
while let Ok(command) = receiver.recv() {
match command {
StderrCommand::Write {
serialized,
completion,
} => {
let stderr = std::io::stderr();
let mut stderr = stderr.lock();
let result = stderr
.write_all(serialized.as_bytes())
.and_then(|_| stderr.write_all(b"\n"))
.and_then(|_| stderr.flush());
if let Err(error) = result {
let thread_error =
std::io::Error::new(error.kind(), error.to_string());
drop(completion.send(Err(error)));
return Err(thread_error);
}
drop(completion.send(Ok(())));
}
StderrCommand::Close => break,
}
}
Ok(())
})?;
Ok(Self {
sender,
worker: Mutex::new(Some(worker)),
})
}
}
#[async_trait]
impl AuditOutput for StderrOutput {
fn name(&self) -> &str {
"stderr"
}
async fn write(&self, event: &AuditEvent) -> AuditResult<()> {
let serialized = event.to_json()?;
let (completion, completed) = tokio::sync::oneshot::channel();
let command = StderrCommand::Write {
serialized,
completion,
};
self.enqueue(command).await?;
completed
.await
.map_err(|_| AuditError::Channel("audit stderr writer stopped".to_string()))??;
Ok(())
}
async fn flush(&self) -> AuditResult<()> {
Ok(())
}
async fn close(&self) -> AuditResult<()> {
let Some(worker) = self.worker.lock().await.take() else {
return Ok(());
};
let sender = self.sender.clone();
tokio::task::spawn_blocking(move || {
let close_result = sender.send(StderrCommand::Close);
let worker_result = worker
.join()
.map_err(|_| AuditError::Output("audit stderr writer panicked".to_string()))?;
worker_result?;
close_result.map_err(|_| {
AuditError::Channel("audit stderr writer stopped before close".to_string())
})
})
.await
.map_err(|error| AuditError::Output(format!("audit stderr join failed: {error}")))?
}
}
impl StderrOutput {
async fn enqueue(&self, command: StderrCommand) -> AuditResult<()> {
match self.sender.try_send(command) {
Ok(()) => Ok(()),
Err(TrySendError::Disconnected(_)) => Err(AuditError::Channel(
"audit stderr writer stopped".to_string(),
)),
Err(TrySendError::Full(command)) => {
let sender = self.sender.clone();
tokio::task::spawn_blocking(move || sender.send(command))
.await
.map_err(|error| {
AuditError::Output(format!("audit stderr enqueue failed: {error}"))
})?
.map_err(|_| AuditError::Channel("audit stderr writer stopped".to_string()))
}
}
}
}
pub struct FileOutput {
path: PathBuf,
file: Arc<Mutex<Option<File>>>,
buffer: Arc<Mutex<Vec<String>>>,
buffer_size: usize,
}
impl FileOutput {
pub async fn new(path: impl Into<PathBuf>) -> AuditResult<Self> {
let path = path.into();
if let Some(parent) = path.parent() {
tokio::fs::create_dir_all(parent).await?;
}
let file = OpenOptions::new()
.create(true)
.append(true)
.open(&path)
.await?;
Ok(Self {
path,
file: Arc::new(Mutex::new(Some(file))),
buffer: Arc::new(Mutex::new(Vec::new())),
buffer_size: 100,
})
}
pub fn with_buffer_size(mut self, size: usize) -> Self {
self.buffer_size = size;
self
}
pub fn path(&self) -> &PathBuf {
&self.path
}
async fn write_buffer(&self) -> AuditResult<()> {
let mut buffer = self.buffer.lock().await;
if buffer.is_empty() {
return Ok(());
}
let mut file_guard = self.file.lock().await;
if let Some(ref mut file) = *file_guard {
for line in buffer.drain(..) {
file.write_all(line.as_bytes()).await?;
file.write_all(b"\n").await?;
}
file.flush().await?;
}
Ok(())
}
}
#[async_trait]
impl AuditOutput for FileOutput {
fn name(&self) -> &str {
"file"
}
async fn write(&self, event: &AuditEvent) -> AuditResult<()> {
let json = event.to_json()?;
let mut buffer = self.buffer.lock().await;
buffer.push(json);
if buffer.len() >= self.buffer_size {
drop(buffer);
self.write_buffer().await?;
}
Ok(())
}
async fn flush(&self) -> AuditResult<()> {
self.write_buffer().await
}
async fn close(&self) -> AuditResult<()> {
self.flush().await?;
let mut file_guard = self.file.lock().await;
*file_guard = None;
Ok(())
}
}
pub struct MemoryOutput {
events: Arc<Mutex<VecDeque<AuditEvent>>>,
max_events: usize,
}
impl MemoryOutput {
pub fn new(max_events: usize) -> Self {
Self {
events: Arc::new(Mutex::new(VecDeque::new())),
max_events,
}
}
pub async fn events(&self) -> Vec<AuditEvent> {
let events = self.events.lock().await;
events.iter().cloned().collect()
}
pub async fn count(&self) -> usize {
let events = self.events.lock().await;
events.len()
}
pub async fn clear(&self) {
let mut events = self.events.lock().await;
events.clear();
}
pub async fn last_n(&self, n: usize) -> Vec<AuditEvent> {
let events = self.events.lock().await;
events.iter().rev().take(n).cloned().collect()
}
}
impl Default for MemoryOutput {
fn default() -> Self {
Self::new(1000)
}
}
#[async_trait]
impl AuditOutput for MemoryOutput {
fn name(&self) -> &str {
"memory"
}
async fn write(&self, event: &AuditEvent) -> AuditResult<()> {
let mut events = self.events.lock().await;
while events.len() >= self.max_events {
events.pop_front();
}
events.push_back(event.clone());
Ok(())
}
async fn flush(&self) -> AuditResult<()> {
Ok(())
}
async fn close(&self) -> AuditResult<()> {
Ok(())
}
}
pub struct NullOutput;
#[async_trait]
impl AuditOutput for NullOutput {
fn name(&self) -> &str {
"null"
}
async fn write(&self, _event: &AuditEvent) -> AuditResult<()> {
Ok(())
}
async fn flush(&self) -> AuditResult<()> {
Ok(())
}
async fn close(&self) -> AuditResult<()> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::audit::events::EventType;
#[tokio::test]
async fn stderr_handoff_applies_bounded_backpressure_without_dropping() {
let (sender, receiver) = std::sync::mpsc::sync_channel(1);
let output = StderrOutput {
sender,
worker: Mutex::new(None),
};
let command = |serialized: &str| {
let (completion, _) = tokio::sync::oneshot::channel();
StderrCommand::Write {
serialized: serialized.to_string(),
completion,
}
};
output.enqueue(command("event-1")).await.unwrap();
let second_write = output.enqueue(command("event-2"));
tokio::pin!(second_write);
assert!(
tokio::time::timeout(std::time::Duration::from_millis(20), &mut second_write)
.await
.is_err()
);
assert!(matches!(
receiver.recv().unwrap(),
StderrCommand::Write { .. }
));
tokio::time::timeout(std::time::Duration::from_secs(1), second_write)
.await
.unwrap()
.unwrap();
assert!(matches!(
receiver.recv().unwrap(),
StderrCommand::Write { .. }
));
}
#[tokio::test]
async fn test_memory_output() {
let output = MemoryOutput::new(10);
for i in 0..5 {
let event = AuditEvent::new(EventType::System, format!("Event {}", i));
output.write(&event).await.unwrap();
}
assert_eq!(output.count().await, 5);
let events = output.events().await;
assert_eq!(events.len(), 5);
}
#[tokio::test]
async fn test_memory_output_max_events() {
let output = MemoryOutput::new(3);
for i in 0..5 {
let event = AuditEvent::new(EventType::System, format!("Event {}", i));
output.write(&event).await.unwrap();
}
assert_eq!(output.count().await, 3);
let events = output.events().await;
assert!(events[0].message.contains("Event 2"));
assert!(events[1].message.contains("Event 3"));
assert!(events[2].message.contains("Event 4"));
}
#[tokio::test]
async fn test_memory_output_clear() {
let output = MemoryOutput::new(10);
let event = AuditEvent::new(EventType::System, "Test");
output.write(&event).await.unwrap();
assert_eq!(output.count().await, 1);
output.clear().await;
assert_eq!(output.count().await, 0);
}
#[tokio::test]
async fn test_memory_output_last_n() {
let output = MemoryOutput::new(10);
for i in 0..5 {
let event = AuditEvent::new(EventType::System, format!("Event {}", i));
output.write(&event).await.unwrap();
}
let last_2 = output.last_n(2).await;
assert_eq!(last_2.len(), 2);
assert!(last_2[0].message.contains("Event 4"));
assert!(last_2[1].message.contains("Event 3"));
}
#[tokio::test]
async fn test_null_output() {
let output = NullOutput;
let event = AuditEvent::new(EventType::System, "Test");
output.write(&event).await.unwrap();
output.flush().await.unwrap();
output.close().await.unwrap();
}
#[tokio::test]
async fn test_file_output() {
let temp_dir = std::env::temp_dir();
let path = temp_dir.join("test_audit.log");
let _ = tokio::fs::remove_file(&path).await;
let output = FileOutput::new(&path).await.unwrap();
let event = AuditEvent::new(EventType::System, "Test event");
output.write(&event).await.unwrap();
output.flush().await.unwrap();
let content = tokio::fs::read_to_string(&path).await.unwrap();
assert!(content.contains("Test event"));
output.close().await.unwrap();
let _ = tokio::fs::remove_file(&path).await;
}
#[tokio::test]
async fn test_file_output_buffering() {
let temp_dir = std::env::temp_dir();
let path = temp_dir.join("test_audit_buffer.log");
let _ = tokio::fs::remove_file(&path).await;
let output = FileOutput::new(&path).await.unwrap().with_buffer_size(3);
for i in 0..2 {
let event = AuditEvent::new(EventType::System, format!("Event {}", i));
output.write(&event).await.unwrap();
}
let content = tokio::fs::read_to_string(&path).await.unwrap();
assert!(content.is_empty());
let event = AuditEvent::new(EventType::System, "Event 2");
output.write(&event).await.unwrap();
let content = tokio::fs::read_to_string(&path).await.unwrap();
assert!(content.contains("Event 0"));
assert!(content.contains("Event 1"));
assert!(content.contains("Event 2"));
output.close().await.unwrap();
let _ = tokio::fs::remove_file(&path).await;
}
}