Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 18 additions & 5 deletions src/source/files.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ use std::fs::{self, File};
use std::io::Read;
use std::path::{Path, PathBuf};

use super::{ProcessStats, Source};
use super::{OutputGuard, ProcessStats, Source};
use crate::derive::KeyDeriver;
use crate::matcher::Matcher;
use crate::output::Output;
Expand Down Expand Up @@ -103,8 +103,13 @@ impl Source for FilesSource {
let processed = std::sync::atomic::AtomicU64::new(0);
let stats = std::sync::atomic::AtomicU64::new(0);
let matches = std::sync::atomic::AtomicU64::new(0);
let guard = OutputGuard::new();

self.files.par_iter().for_each(|path| {
if guard.is_poisoned() {
Comment thread
oritwoen marked this conversation as resolved.
return;
}

let contents = match read_file_contents(path) {
Ok(c) => c,
Err(e) => {
Expand All @@ -125,17 +130,24 @@ impl Source for FilesSource {
transform.apply_batch(&inputs, &mut buffer);

for (source, key) in &buffer {
if guard.is_poisoned() {
break;
}
Comment thread
oritwoen marked this conversation as resolved.

let derived = deriver.derive(key);

if let Some(m) = matcher {
if let Some(match_info) = m.check(&derived) {
output
.hit(source, transform.name(), &derived, &match_info)
.ok();
guard.check(output.hit(
source,
transform.name(),
&derived,
&match_info,
));
matches.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
} else {
output.key(source, transform.name(), &derived).ok();
guard.check(output.key(source, transform.name(), &derived));
}

stats.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Expand All @@ -146,6 +158,7 @@ impl Source for FilesSource {
});

pb.finish_and_clear();
guard.into_result()?;

Ok(ProcessStats {
inputs_processed: processed.load(std::sync::atomic::Ordering::Relaxed),
Expand Down
72 changes: 72 additions & 0 deletions src/source/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -48,3 +48,75 @@ pub enum SourceType {
Timestamps,
Stdin,
}

/// Captures output errors inside Rayon closures where `?` can't be used.
///
/// Check `is_poisoned()` before each output call to skip work after failure.
/// Call `into_result()` after the parallel section to propagate the first error.
pub(crate) struct OutputGuard {
poisoned: std::sync::atomic::AtomicBool,
first_error: std::sync::Mutex<Option<String>>,
}

impl OutputGuard {
pub fn new() -> Self {
Self {
poisoned: std::sync::atomic::AtomicBool::new(false),
first_error: std::sync::Mutex::new(None),
}
}

pub fn is_poisoned(&self) -> bool {
self.poisoned.load(std::sync::atomic::Ordering::Relaxed)
}

pub fn check(&self, result: Result<()>) {
if let Err(e) = result {
self.poisoned
.store(true, std::sync::atomic::Ordering::Relaxed);
if let Ok(mut first) = self.first_error.lock() {
if first.is_none() {
*first = Some(e.to_string());
}
}
}
}

pub fn into_result(self) -> Result<()> {
if self.poisoned.load(std::sync::atomic::Ordering::Relaxed) {
let msg = self
.first_error
.into_inner()
.unwrap_or_else(|e| e.into_inner())
.unwrap_or_else(|| "unknown output error".to_string());
Comment thread
oritwoen marked this conversation as resolved.
anyhow::bail!("Output failed: {}", msg)
} else {
Ok(())
}
}
}

#[cfg(test)]
mod tests {
use super::*;

#[test]
fn output_guard_ok_stays_clean() {
let guard = OutputGuard::new();
guard.check(Ok(()));
guard.check(Ok(()));
assert!(!guard.is_poisoned());
assert!(guard.into_result().is_ok());
}

#[test]
fn output_guard_captures_first_error() {
let guard = OutputGuard::new();
guard.check(Ok(()));
guard.check(Err(anyhow::anyhow!("disk full")));
guard.check(Err(anyhow::anyhow!("second error")));
assert!(guard.is_poisoned());
let err = guard.into_result().unwrap_err();
assert!(err.to_string().contains("disk full"));
}
}
62 changes: 56 additions & 6 deletions src/source/range.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ use anyhow::Result;
use indicatif::ProgressBar;
use rayon::prelude::*;

use super::{ProcessStats, Source};
use super::{OutputGuard, ProcessStats, Source};
use crate::derive::KeyDeriver;
use crate::matcher::Matcher;
use crate::output::Output;
Expand Down Expand Up @@ -52,10 +52,15 @@ impl Source for RangeSource {

let stats = std::sync::atomic::AtomicU64::new(0);
let matches = std::sync::atomic::AtomicU64::new(0);
let guard = OutputGuard::new();

let num_batches = count / BATCH_SIZE + u64::from(count % BATCH_SIZE != 0);

(0..num_batches).into_par_iter().for_each(|batch_idx| {
if guard.is_poisoned() {
Comment thread
oritwoen marked this conversation as resolved.
return;
}

let batch_start = self.start + batch_idx * BATCH_SIZE;
let batch_end = batch_start.saturating_add(BATCH_SIZE - 1).min(self.end);

Comment thread
oritwoen marked this conversation as resolved.
Expand All @@ -67,17 +72,24 @@ impl Source for RangeSource {
transform.apply_batch(&inputs, &mut buffer);

for (source, key) in &buffer {
if guard.is_poisoned() {
break;
Comment thread
oritwoen marked this conversation as resolved.
}

let derived = deriver.derive(key);

if let Some(m) = matcher {
if let Some(match_info) = m.check(&derived) {
output
.hit(source, transform.name(), &derived, &match_info)
.ok();
guard.check(output.hit(
source,
transform.name(),
&derived,
&match_info,
Comment thread
oritwoen marked this conversation as resolved.
));
matches.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
} else {
output.key(source, transform.name(), &derived).ok();
guard.check(output.key(source, transform.name(), &derived));
}

stats.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Expand All @@ -88,6 +100,7 @@ impl Source for RangeSource {
});

pb.finish_and_clear();
guard.into_result()?;

Ok(ProcessStats {
inputs_processed: count,
Expand All @@ -101,7 +114,44 @@ impl Source for RangeSource {
mod tests {
use super::*;
use crate::derive::KeyDeriver;
use crate::output::ConsoleOutput;
use crate::matcher::MatchInfo;
use crate::output::{ConsoleOutput, Output};
use crate::transform::TransformType;

struct FailingOutput;

impl Output for FailingOutput {
fn key(&self, _: &str, _: &str, _: &crate::derive::DerivedKey) -> anyhow::Result<()> {
anyhow::bail!("broken pipe")
}
fn hit(
&self,
_: &str,
_: &str,
_: &crate::derive::DerivedKey,
_: &MatchInfo,
) -> anyhow::Result<()> {
anyhow::bail!("broken pipe")
}
fn flush(&self) -> anyhow::Result<()> {
Ok(())
}
}

#[test]
fn process_propagates_output_error() {
let source = RangeSource::new(1, 10);
let deriver = KeyDeriver::new();
let output = FailingOutput;
let transforms: Vec<Box<dyn Transform>> = vec![TransformType::Sha256.create()];

let result = source.process(&transforms, &deriver, None, &output);
assert!(result.is_err());
assert!(
result.unwrap_err().to_string().contains("broken pipe"),
"error message should contain the original output error"
);
}

#[test]
fn process_rejects_descending_range() {
Expand Down
15 changes: 7 additions & 8 deletions src/source/stdin.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
use anyhow::Result;
use std::io::{self, BufRead};


use super::{ProcessStats, Source};
use crate::derive::KeyDeriver;
use crate::matcher::Matcher;
Expand Down Expand Up @@ -56,7 +57,7 @@ impl Source for StdinSource {
// Process in batches
if batch.len() >= 1000 {
let (keys, found) =
process_batch(&batch, transforms, &deriver, matcher, output, &mut buffer);
process_batch(&batch, transforms, &deriver, matcher, output, &mut buffer)?;
keys_generated += keys;
matches_found += found;
batch.clear();
Expand All @@ -66,7 +67,7 @@ impl Source for StdinSource {
// Process remaining
if !batch.is_empty() {
let (keys, found) =
process_batch(&batch, transforms, &deriver, matcher, output, &mut buffer);
process_batch(&batch, transforms, &deriver, matcher, output, &mut buffer)?;
keys_generated += keys;
matches_found += found;
}
Expand All @@ -86,7 +87,7 @@ fn process_batch(
matcher: Option<&Matcher>,
output: &dyn Output,
buffer: &mut Vec<(String, [u8; 32])>,
) -> (u64, u64) {
) -> Result<(u64, u64)> {
let mut keys_generated = 0u64;
let mut matches_found = 0u64;

Expand All @@ -99,18 +100,16 @@ fn process_batch(

if let Some(m) = matcher {
if let Some(match_info) = m.check(&derived) {
output
.hit(source, transform.name(), &derived, &match_info)
.ok();
output.hit(source, transform.name(), &derived, &match_info)?;
matches_found += 1;
}
} else {
output.key(source, transform.name(), &derived).ok();
output.key(source, transform.name(), &derived)?;
}

keys_generated += 1;
}
}

(keys_generated, matches_found)
Ok((keys_generated, matches_found))
}
28 changes: 21 additions & 7 deletions src/source/timestamps.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ use chrono::NaiveDate;
use indicatif::ProgressBar;
use rayon::prelude::*;

use super::{ProcessStats, Source};
use super::{OutputGuard, ProcessStats, Source};
use crate::derive::KeyDeriver;
use crate::matcher::Matcher;
use crate::output::Output;
Expand Down Expand Up @@ -105,17 +105,27 @@ impl Source for TimestampSource {

let stats = std::sync::atomic::AtomicU64::new(0);
let matches = std::sync::atomic::AtomicU64::new(0);
let guard = OutputGuard::new();

(self.start..=self.end).into_par_iter().for_each(|ts| {
if guard.is_poisoned() {
return;
}

// Process base timestamp
process_timestamp(ts, transforms, &deriver, matcher, output, &stats, &matches);
process_timestamp(
ts, transforms, &deriver, matcher, output, &guard, &stats, &matches,
);

// Process milliseconds if enabled
if self.milliseconds {
for ms in 0u64..1000 {
if guard.is_poisoned() {
return;
}
let ts_ms = ts * 1000 + ms;
process_timestamp(
ts_ms, transforms, &deriver, matcher, output, &stats, &matches,
ts_ms, transforms, &deriver, matcher, output, &guard, &stats, &matches,
);
}
pb.inc(1001);
Expand All @@ -125,6 +135,7 @@ impl Source for TimestampSource {
});

pb.finish_and_clear();
guard.into_result()?;

Ok(ProcessStats {
inputs_processed: total,
Expand Down Expand Up @@ -207,6 +218,7 @@ fn process_timestamp(
deriver: &KeyDeriver,
matcher: Option<&Matcher>,
output: &dyn Output,
guard: &OutputGuard,
stats: &std::sync::atomic::AtomicU64,
matches: &std::sync::atomic::AtomicU64,
) {
Expand All @@ -218,17 +230,19 @@ fn process_timestamp(
transform.apply_batch(&inputs, &mut buffer);

for (source, key) in &buffer {
if guard.is_poisoned() {
break;
}
Comment thread
oritwoen marked this conversation as resolved.

let derived = deriver.derive(key);

if let Some(m) = matcher {
if let Some(match_info) = m.check(&derived) {
output
.hit(source, transform.name(), &derived, &match_info)
.ok();
guard.check(output.hit(source, transform.name(), &derived, &match_info));
Comment thread
oritwoen marked this conversation as resolved.
matches.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
} else {
output.key(source, transform.name(), &derived).ok();
guard.check(output.key(source, transform.name(), &derived));
}

stats.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Expand Down
Loading
Loading