use super::{
BatchKey, FailureReporter, RunnerConfig, ValidationOutcome, batch_level, finalize_result, session_for, session_key,
session_lock_for, session_preparation_error, session_preparation_result,
};
use crate::snippets::error::Result;
use crate::snippets::session::ValidationSession;
use crate::snippets::types::{Snippet, SnippetStatus, ValidationLevel, ValidationResult};
use crate::snippets::validators::{BatchValidation, SnippetValidator, ValidatorRegistry};
use rayon::prelude::*;
use std::collections::{BTreeMap, HashMap};
use std::sync::Mutex;
use std::time::Instant;
struct BatchContext<'a> {
snippets: &'a [Snippet],
registry: &'a ValidatorRegistry,
config: &'a RunnerConfig,
sessions: &'a HashMap<String, ValidationSession>,
session_locks: &'a HashMap<String, Mutex<()>>,
reporter: &'a FailureReporter,
}
struct GroupedSnippets {
groups: BTreeMap<BatchKey, Vec<usize>>,
results: Vec<Option<ValidationResult>>,
}
pub(super) fn validate_batches(
snippets: &[Snippet],
registry: &ValidatorRegistry,
config: &RunnerConfig,
sessions: &HashMap<String, ValidationSession>,
session_errors: &HashMap<String, crate::snippets::session::SessionPreparationError>,
session_locks: &HashMap<String, Mutex<()>>,
reporter: &FailureReporter,
) -> Vec<Option<ValidationResult>> {
let GroupedSnippets { groups, mut results } =
group_batchable_snippets(snippets, registry, config, sessions, session_errors, reporter);
let context = BatchContext {
snippets,
registry,
config,
sessions,
session_locks,
reporter,
};
for (index, validated) in dispatch_groups(&context, groups) {
results[index] = Some(validated);
}
results
}
fn group_batchable_snippets(
snippets: &[Snippet],
registry: &ValidatorRegistry,
config: &RunnerConfig,
sessions: &HashMap<String, ValidationSession>,
session_errors: &HashMap<String, crate::snippets::session::SessionPreparationError>,
reporter: &FailureReporter,
) -> GroupedSnippets {
let mut results = vec![None; snippets.len()];
let mut groups = BTreeMap::<BatchKey, Vec<usize>>::new();
for (index, snippet) in snippets.iter().enumerate() {
if let Some(preparation_error) = session_preparation_error(snippet, config, session_errors) {
let failure = session_preparation_result(snippet, config, &preparation_error);
reporter.record(&failure);
results[index] = Some(failure);
continue;
}
let session = session_for(snippet, sessions);
if let Some(level) = batch_level(snippet, registry, config, session) {
let key = (
snippet.language,
session_key(snippet, sessions).map(str::to_string),
level,
);
groups.entry(key).or_default().push(index);
}
}
GroupedSnippets { groups, results }
}
fn dispatch_groups(
context: &BatchContext<'_>,
groups: BTreeMap<BatchKey, Vec<usize>>,
) -> Vec<(usize, ValidationResult)> {
let span = tracing::Span::current();
groups
.into_iter()
.collect::<Vec<_>>()
.into_par_iter()
.map(|(key, indices)| span.in_scope(|| validate_group(context, &key, &indices)))
.collect::<Vec<_>>()
.into_iter()
.flatten()
.collect()
}
const BATCH_TIMEOUT_SNIPPETS_PER_BUDGET: u64 = 8;
fn batch_timeout_secs(per_snippet_secs: u64, snippet_count: usize) -> u64 {
let count = u64::try_from(snippet_count).unwrap_or(u64::MAX);
let grants = count.div_ceil(BATCH_TIMEOUT_SNIPPETS_PER_BUDGET).max(1);
per_snippet_secs.saturating_mul(grants)
}
fn validate_group(context: &BatchContext<'_>, key: &BatchKey, indices: &[usize]) -> Vec<(usize, ValidationResult)> {
let (language, session_target, level) = key;
let validator = context.registry.get(*language).expect("batch group validator");
let session = session_target.as_deref().and_then(|value| context.sessions.get(value));
let batch_snippets = indices
.iter()
.map(|index| &context.snippets[*index])
.collect::<Vec<_>>();
let timeout_secs = batch_timeout_secs(context.config.timeout_secs, batch_snippets.len());
tracing::info!(
language = %language,
snippet_count = batch_snippets.len(),
timeout_secs,
"Starting batched snippet validation"
);
let started = Instant::now();
let batch = run_batch(context, validator, session, *level, &batch_snippets, timeout_secs);
let Some(batch) = batch else {
tracing::info!(
language = %language,
snippet_count = batch_snippets.len(),
"Batch validation declined for this group; falling back to per-snippet validation"
);
return Vec::new();
};
let values = batch_statuses(batch, indices.len());
let duration_ms = u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX);
tracing::info!(
language = %language,
snippet_count = batch_snippets.len(),
duration_ms,
"Finished batched snippet validation"
);
finalize_group(
context,
validator,
session,
*level,
&batch_snippets,
values,
indices,
duration_ms,
)
}
fn run_batch(
context: &BatchContext<'_>,
validator: &dyn SnippetValidator,
session: Option<&ValidationSession>,
level: ValidationLevel,
batch_snippets: &[&Snippet],
timeout_secs: u64,
) -> Option<Result<BatchValidation>> {
let validation = || validator.validate_batch_in_session(batch_snippets, level, timeout_secs, session);
match session_lock_for(session, context.session_locks) {
Some(lock) => lock.lock().ok().and_then(|_guard| validation()),
None => validation(),
}
}
fn batch_statuses(batch: Result<BatchValidation>, expected: usize) -> BatchValidation {
match batch {
Ok(values) if values.len() == expected => values,
Ok(values) => {
let message = format!(
"batch validator returned {} results for {expected} snippets",
values.len()
);
vec![(SnippetStatus::Error, Some(message)); expected]
}
Err(error) => vec![(SnippetStatus::Error, Some(error.to_string())); expected],
}
}
#[expect(
clippy::too_many_arguments,
reason = "one call site; splitting it further would only move the arguments"
)]
fn finalize_group(
context: &BatchContext<'_>,
validator: &dyn SnippetValidator,
session: Option<&ValidationSession>,
level: ValidationLevel,
batch_snippets: &[&Snippet],
values: BatchValidation,
indices: &[usize],
duration_ms: u64,
) -> Vec<(usize, ValidationResult)> {
let mut finalized = Vec::with_capacity(indices.len());
for ((index, snippet), (status, message)) in indices.iter().copied().zip(batch_snippets).zip(values) {
let value = finalize_result(
snippet,
validator,
context.config,
session,
level,
ValidationOutcome {
status,
message,
duration_ms,
},
);
context.reporter.record(&value);
finalized.push((index, value));
}
finalized
}
#[cfg(test)]
mod tests {
use super::batch_timeout_secs;
#[test]
fn a_batch_budget_grows_with_the_number_of_snippets_it_covers() {
assert_eq!(batch_timeout_secs(120, 1), 120, "a single snippet gets one budget");
assert_eq!(
batch_timeout_secs(120, 8),
120,
"a batch under the divisor still gets one"
);
assert_eq!(
batch_timeout_secs(120, 9),
240,
"one snippet past the divisor buys the next grant"
);
assert_eq!(
batch_timeout_secs(120, 283),
4320,
"a full language's batch gets 36 grants"
);
}
#[test]
fn an_empty_batch_still_gets_one_whole_budget() {
assert_eq!(batch_timeout_secs(120, 0), 120);
}
#[test]
fn a_batch_budget_stays_far_below_the_serial_path_it_replaces() {
let serial = 120 * 283;
let granted = batch_timeout_secs(120, 283);
assert!(
granted <= serial / 4,
"a 283-snippet batch was granted {granted}s against the {serial}s the serial path would spend; \
the divisor is what keeps a hung compiler from running for hours before the timeout fires"
);
}
use crate::snippets::runner::{RunnerConfig, run_validation};
use crate::snippets::types::{
Language, Snippet, SnippetMetadata, SnippetStatus, SourceOrigin, ValidationLevel, ValidationResult,
};
use crate::snippets::validators::{SnippetValidator, ValidatorRegistry};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::{Duration, Instant};
const CONCURRENCY_PROBE_TIMEOUT: Duration = Duration::from_secs(10);
#[derive(Default)]
struct ConcurrencyProbe {
in_flight: AtomicUsize,
peak: AtomicUsize,
}
struct ProbingBatchValidator {
language: Language,
probe: Arc<ConcurrencyProbe>,
}
impl SnippetValidator for ProbingBatchValidator {
fn language(&self) -> Language {
self.language
}
fn is_available(&self) -> bool {
true
}
fn validate(
&self,
_snippet: &Snippet,
_level: ValidationLevel,
_timeout_secs: u64,
) -> crate::snippets::error::Result<(SnippetStatus, Option<String>)> {
Ok((SnippetStatus::Pass, None))
}
fn validate_batch_in_session(
&self,
snippets: &[&Snippet],
_level: ValidationLevel,
_timeout_secs: u64,
_session: Option<&crate::snippets::session::ValidationSession>,
) -> Option<crate::snippets::error::Result<Vec<(SnippetStatus, Option<String>)>>> {
let entered = self.probe.in_flight.fetch_add(1, Ordering::SeqCst) + 1;
self.probe.peak.fetch_max(entered, Ordering::SeqCst);
let deadline = Instant::now() + CONCURRENCY_PROBE_TIMEOUT;
while self.probe.peak.load(Ordering::SeqCst) < 2 && Instant::now() < deadline {
std::thread::sleep(Duration::from_millis(1));
}
self.probe.in_flight.fetch_sub(1, Ordering::SeqCst);
let language = self.language;
Some(Ok(snippets
.iter()
.map(|_| (SnippetStatus::Fail, Some(format!("{language} batch"))))
.collect()))
}
fn supports_batching(&self) -> bool {
true
}
fn max_level(&self) -> ValidationLevel {
ValidationLevel::Run
}
fn is_dependency_error(&self, _output: &str) -> bool {
false
}
}
fn snippet(language: Language) -> Snippet {
Snippet {
id: None,
path: "example.md".into(),
language,
title: None,
code: "example".into(),
start_line: 1,
block_index: 0,
annotation: None,
metadata: SnippetMetadata::default(),
source_origin: SourceOrigin {
path: "example.md".into(),
line: 1,
block_index: 0,
},
}
}
fn probing_registry(probe: &Arc<ConcurrencyProbe>) -> ValidatorRegistry {
let mut registry = ValidatorRegistry::new();
for language in [Language::Rust, Language::Java] {
registry.register(Box::new(ProbingBatchValidator {
language,
probe: Arc::clone(probe),
}));
}
registry
}
fn batch_config() -> RunnerConfig {
RunnerConfig {
level: ValidationLevel::Compile,
parallelism: 4,
cache_dir: None,
..RunnerConfig::default()
}
}
#[test]
fn batch_groups_with_different_keys_run_concurrently() {
let probe = Arc::new(ConcurrencyProbe::default());
let registry = probing_registry(&probe);
let snippets = [snippet(Language::Rust), snippet(Language::Java)];
let summary = run_validation(&snippets, ®istry, &batch_config()).expect("validation completes");
assert_eq!(summary.results.len(), 2);
assert_eq!(probe.peak.load(Ordering::SeqCst), 2);
}
#[test]
fn concurrent_groups_keep_results_at_their_snippet_positions() {
let probe = Arc::new(ConcurrencyProbe::default());
let registry = probing_registry(&probe);
let snippets = [
snippet(Language::Rust),
snippet(Language::Java),
snippet(Language::Java),
snippet(Language::Rust),
];
let summary = run_validation(&snippets, ®istry, &batch_config()).expect("validation completes");
let messages = summary
.results
.iter()
.map(|result: &ValidationResult| result.message.clone().unwrap_or_default())
.collect::<Vec<_>>();
assert_eq!(
messages,
vec![
"rust batch".to_string(),
"java batch".to_string(),
"java batch".to_string(),
"rust batch".to_string(),
]
);
assert_eq!(summary.failed, 4);
}
}