use crate::error::{Result, VmSpectError};
use crate::models::options::{InspectionProgress, InspectionProgressEvent, Options};
use crate::models::traits::AnalysisResult;
use crate::models::InspectionReport;
use crate::parsers;
use crate::vms;
use crate::vms::nbd::{self, NbdReader};
use crate::vms::stream::{identify_image, DiskReader};
use std::collections::VecDeque;
use std::path::Path;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::thread::JoinHandle;
use std::time::Instant;
pub struct ConcurrentProcessor;
impl ConcurrentProcessor {
pub fn process_in_parallel<T, R, F>(
items: Vec<T>,
cancel_token: Option<Arc<AtomicBool>>,
progress: Option<Arc<InspectionProgress>>,
max_workers: usize,
f: F,
) -> Result<Vec<R>>
where
T: Send + 'static,
R: Send + 'static,
F: Fn(T) -> Result<R> + Send + Sync + 'static,
{
if items.is_empty() {
return Ok(Vec::new());
}
if let Some(ref cancel) = cancel_token {
if cancel.load(Ordering::Acquire) {
return Ok(Vec::new());
}
}
if let Some(ref p) = progress {
p.set_total_tasks(items.len());
}
let total_items = items.len();
let num_workers = max_workers.max(1).min(total_items).min(32);
let queue = Arc::new(Mutex::new(
items
.into_iter()
.enumerate()
.collect::<VecDeque<(usize, T)>>(),
));
let results = Arc::new(Mutex::new(Vec::<(usize, R)>::with_capacity(total_items)));
let stored_error = Arc::new(Mutex::new(None::<VmSpectError>));
let f = Arc::new(f);
let mut handles: Vec<JoinHandle<()>> = Vec::with_capacity(num_workers);
for worker_id in 0..num_workers {
let queue_clone = Arc::clone(&queue);
let results_clone = Arc::clone(&results);
let error_clone = Arc::clone(&stored_error);
let cancel_clone = cancel_token.clone();
let progress_clone = progress.clone();
let f_clone = Arc::clone(&f);
let builder = std::thread::Builder::new().name(format!("vmspect-worker-{}", worker_id));
let handle = builder.spawn(move || {
loop {
if let Some(ref cancel) = cancel_clone {
if cancel.load(Ordering::Acquire) {
break;
}
}
let task = {
let mut q = queue_clone.lock().unwrap_or_else(|e| e.into_inner());
q.pop_front()
};
let Some((idx, item)) = task else {
break;
};
if let Some(ref cancel) = cancel_clone {
if cancel.load(Ordering::Acquire) {
break;
}
}
let result = f_clone(item);
match result {
Ok(value) => {
let mut res = results_clone.lock().unwrap_or_else(|e| e.into_inner());
res.push((idx, value));
if let Some(ref p) = progress_clone {
p.increment_completed_tasks();
}
}
Err(e) => {
if !matches!(e, VmSpectError::Cancelled) {
let mut err_guard =
error_clone.lock().unwrap_or_else(|e| e.into_inner());
if err_guard.is_none() {
*err_guard = Some(e);
}
}
break;
}
}
}
});
if let Ok(h) = handle {
handles.push(h);
}
}
for handle in handles {
let _ = handle.join();
}
let was_cancelled = cancel_token
.as_ref()
.map(|c| c.load(Ordering::Acquire))
.unwrap_or(false);
if !was_cancelled {
if let Some(err) = stored_error
.lock()
.unwrap_or_else(|e| e.into_inner())
.take()
{
return Err(err);
}
}
let mut res = Arc::try_unwrap(results)
.map(|m| m.into_inner().unwrap_or_else(|e| e.into_inner()))
.unwrap_or_else(|m| std::mem::take(&mut *m.lock().unwrap_or_else(|e| e.into_inner())));
res.sort_by_key(|(idx, _)| *idx);
Ok(res.into_iter().map(|(_, val)| val).collect())
}
pub fn inspect_images<P: AsRef<Path> + Send + 'static>(
paths: Vec<P>,
options: &Options,
max_workers: usize,
) -> Result<Vec<InspectionReport>> {
let options = options.clone();
let cancel = options.cancel_token.clone();
Self::process_in_parallel(paths, cancel, None, max_workers, move |path| {
let engine = InspectionEngine::new(options.clone());
engine.inspect(path.as_ref())
})
}
}
#[derive(Debug, Clone)]
pub struct InspectionEngine {
options: Options,
progress: Arc<InspectionProgress>,
cancel_token: Arc<AtomicBool>,
}
impl Default for InspectionEngine {
fn default() -> Self {
Self::new(Options::default())
}
}
impl InspectionEngine {
pub fn new(options: Options) -> Self {
let cancel_token = options
.cancel_token
.clone()
.unwrap_or_else(|| Arc::new(AtomicBool::new(false)));
let progress = Arc::new(InspectionProgress::with_cancellation_token(Some(
&cancel_token,
)));
let mut options = options;
options.cancel_token = Some(cancel_token.clone());
Self {
options,
progress,
cancel_token,
}
}
pub fn with_options(options: Options) -> Self {
Self::new(options)
}
pub fn progress(&self) -> Arc<InspectionProgress> {
self.progress.clone()
}
pub fn completion_percentage(&self) -> f32 {
self.progress.completion_percentage()
}
pub fn cancel(&self) {
self.cancel_token.store(true, Ordering::Release);
self.progress.cancel();
}
pub fn is_cancelled(&self) -> bool {
self.cancel_token.load(Ordering::Acquire) || self.progress.is_cancelled()
}
pub fn inspect(&self, image_path: &Path) -> Result<InspectionReport> {
self.run_inspection(image_path, None)
}
pub fn inspect_with_progress<F>(
&self,
image_path: &Path,
mut callback: F,
) -> Result<InspectionReport>
where
F: FnMut(InspectionProgressEvent),
{
self.run_inspection(image_path, Some(&mut callback))
}
pub fn inspect_background(
&self,
image_path: &Path,
) -> Result<std::thread::JoinHandle<Result<InspectionReport>>> {
let engine = self.clone();
let path = image_path.to_path_buf();
std::thread::Builder::new()
.name("vmspect-bg-inspect".to_string())
.spawn(move || engine.inspect(&path))
.map_err(VmSpectError::Io)
}
pub fn inspect_batch<P: AsRef<Path> + Send + 'static>(
&self,
paths: Vec<P>,
max_workers: usize,
) -> Result<Vec<InspectionReport>> {
let options = self.options.clone();
let cancel = Some(self.cancel_token.clone());
let progress = Some(self.progress.clone());
ConcurrentProcessor::process_in_parallel(
paths,
cancel,
progress,
max_workers,
move |path| {
let engine = InspectionEngine::new(options.clone());
engine.inspect(path.as_ref())
},
)
}
fn run_inspection(
&self,
image_path: &Path,
mut callback: Option<&mut dyn FnMut(InspectionProgressEvent)>,
) -> Result<InspectionReport> {
let start = Instant::now();
if !image_path.exists() {
return Err(VmSpectError::ImageNotFound(
image_path.display().to_string(),
));
}
if self.is_cancelled() {
return Err(VmSpectError::Cancelled);
}
let mut effective_options = self.options.clone();
effective_options.cancel_token = Some(self.cancel_token.clone());
self.progress.set_stage_id(1);
self.progress.set_percentage(5);
if let Some(ref mut cb) = callback {
cb(InspectionProgressEvent {
percentage: 5,
stage: "Identifying disk image".into(),
detail: Some(format!("Analyzing {}", image_path.display())),
});
}
let image = identify_image(effective_options.qemu_nbd.as_deref(), image_path)?;
if self.is_cancelled() {
return Err(VmSpectError::Cancelled);
}
self.progress.set_total_bytes(image.virtual_size);
self.progress.set_percentage(15);
if let Some(ref mut cb) = callback {
cb(InspectionProgressEvent {
percentage: 15,
stage: "Initializing read backend".into(),
detail: Some(format!(
"Format: {} | Hypervisor: {} | Size: {}",
image.format,
image.hypervisor.name(),
crate::models::format_bytes(image.virtual_size)
)),
});
}
let reader = if effective_options.force_nbd {
let nbd_path = match nbd::resolve_qemu_nbd(effective_options.qemu_nbd.as_deref()) {
Ok(r) => r,
Err(_) => {
return Err(VmSpectError::QemuNotFound(
"qemu-nbd executable was not found on the system".to_string(),
));
}
};
let nbd_reader = NbdReader::open_with_options(&nbd_path, &image, &effective_options)
.map_err(|e| match e.kind() {
std::io::ErrorKind::Interrupted => VmSpectError::Cancelled,
_ => VmSpectError::Nbd(e.to_string()),
})?;
DiskReader::from_nbd(nbd_reader, &image, Some(self.cancel_token.clone()))
} else {
match DiskReader::open_with_options(&image, &effective_options) {
Ok(r) => r,
Err(e) => {
if e.kind() == std::io::ErrorKind::NotFound {
return Err(VmSpectError::QemuNotFound(
"qemu-nbd executable was not found on the system".to_string(),
));
}
if e.kind() == std::io::ErrorKind::Interrupted {
return Err(VmSpectError::Cancelled);
}
return Err(VmSpectError::Io(e));
}
}
};
let chunk_size = effective_options
.chunk_size
.unwrap_or_else(|| reader.recommended_chunk_size());
if self.is_cancelled() {
return Err(VmSpectError::Cancelled);
}
self.progress.set_stage_id(2);
self.progress.set_percentage(25);
if let Some(ref mut cb) = callback {
cb(InspectionProgressEvent {
percentage: 25,
stage: "Reading partition table".into(),
detail: Some(format!(
"Access: {} | Chunk size: {}",
reader.access_mode(),
crate::models::format_bytes(chunk_size)
)),
});
}
let disk = vms::detector::detect_with_progress(
&reader,
Some(self.cancel_token.clone()),
Some(self.progress.clone()),
)?;
if self.is_cancelled() {
return Err(VmSpectError::Cancelled);
}
self.progress.set_percentage(45);
if let Some(ref mut cb) = callback {
cb(InspectionProgressEvent {
percentage: 45,
stage: "Analyzing file systems".into(),
detail: Some(format!(
"{} partitions found. OS detected: {:?}",
disk.partitions.len(),
disk.operating_system
)),
});
}
self.progress.set_stage_id(3);
self.progress.set_percentage(55);
if let Some(ref mut cb) = callback {
cb(InspectionProgressEvent {
percentage: 55,
stage: format!("Analyzing operating system ({:?})", disk.operating_system),
detail: Some("Starting system files / Registry scan".into()),
});
}
let result = if effective_options.should_analyze_system()
|| effective_options.should_analyze_apps()
{
let inspector = parsers::get_inspector(&disk.operating_system);
match inspector.analyze(&reader, &disk.partitions, chunk_size, &effective_options) {
Ok(r) => r,
Err(e) => {
let msg = format!(
"Could not complete the guest OS analysis; continuing with image/partition data only: {}",
e
);
tracing::warn!("{}", msg);
AnalysisResult {
warnings: vec![msg],
..AnalysisResult::default()
}
}
}
} else {
AnalysisResult::default()
};
if self.is_cancelled() {
return Err(VmSpectError::Cancelled);
}
self.progress.set_stage_id(4);
self.progress.set_percentage(90);
if let Some(ref mut cb) = callback {
cb(InspectionProgressEvent {
percentage: 90,
stage: "Generating final report".into(),
detail: Some(format!(
"{} programs/packages identified",
result.programs.len()
)),
});
}
let mut stats = reader.stats();
stats.duration_ms = start.elapsed().as_millis() as u64;
let report = InspectionReport {
image,
scheme: disk.scheme,
partitions: disk.partitions,
operating_system: disk.operating_system,
guest_info: result.guest_info,
installed_programs: result.programs,
warnings: result.warnings,
stats,
};
self.progress.set_percentage(100);
if let Some(ref mut cb) = callback {
cb(InspectionProgressEvent {
percentage: 100,
stage: "Analysis completed successfully".into(),
detail: None,
});
}
Ok(report)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicUsize;
use std::thread::sleep;
use std::time::Duration;
#[test]
fn test_atomic_progress_metadata_and_percentage() {
let prog = InspectionProgress::new();
assert_eq!(prog.completion_percentage(), 0.0);
assert!(!prog.is_cancelled());
prog.set_percentage(50);
assert_eq!(prog.completion_percentage(), 50.0);
prog.set_total_tasks(10);
prog.increment_completed_tasks();
prog.increment_completed_tasks();
assert_eq!(prog.completed_tasks(), 2);
assert_eq!(prog.total_tasks(), 10);
prog.add_bytes_processed(2048);
assert_eq!(prog.bytes_processed(), 2048);
prog.set_total_bytes(4096);
assert_eq!(prog.total_bytes(), 4096);
prog.set_stage_id(3);
assert_eq!(prog.stage_id(), 3);
let snap = prog.snapshot();
assert_eq!(snap.percentage, 50);
assert_eq!(snap.stage_id, 3);
assert_eq!(snap.completed_tasks, 2);
assert_eq!(snap.total_tasks, 10);
assert_eq!(snap.bytes_processed, 2048);
assert_eq!(snap.total_bytes, 4096);
assert!(!snap.cancelled);
prog.cancel();
assert!(prog.is_cancelled());
assert!(prog.snapshot().cancelled);
}
#[test]
fn test_concurrent_processor_normal_execution() {
let items: Vec<u32> = (1..=20).collect();
let cancel = Arc::new(AtomicBool::new(false));
let prog = Arc::new(InspectionProgress::new());
let results = ConcurrentProcessor::process_in_parallel(
items.clone(),
Some(cancel),
Some(prog.clone()),
4,
|x| Ok(x * 2),
)
.expect("parallel processing successful");
assert_eq!(results.len(), 20);
for (i, &val) in results.iter().enumerate() {
assert_eq!(val, (i as u32 + 1) * 2);
}
assert_eq!(prog.completed_tasks(), 20);
}
#[test]
fn test_concurrent_processor_graceful_shutdown_cancellation() {
let items: Vec<u32> = (1..=50).collect();
let cancel = Arc::new(AtomicBool::new(false));
let prog = Arc::new(InspectionProgress::new());
let cancel_clone = cancel.clone();
let prog_worker = prog.clone();
let tasks_started = Arc::new(AtomicUsize::new(0));
let tasks_started_clone = tasks_started.clone();
let handle = std::thread::spawn(move || {
ConcurrentProcessor::process_in_parallel(
items,
Some(cancel_clone),
Some(prog_worker),
4,
move |_item| {
let num = tasks_started_clone.fetch_add(1, Ordering::SeqCst);
if num >= 2 {
sleep(Duration::from_millis(50));
}
Ok(())
},
)
});
let wait_start = Instant::now();
while prog.completed_tasks() < 2 && wait_start.elapsed() < Duration::from_secs(5) {
sleep(Duration::from_millis(1));
}
cancel.store(true, Ordering::Release);
let result = handle.join().expect("coordinator thread finished");
let partial_results = result.expect("must preserve partial results");
assert!(
!partial_results.is_empty(),
"Must preserve the completed results"
);
assert!(
partial_results.len() < 50,
"Must not process all items if cancelled"
);
let total_started = tasks_started.load(Ordering::SeqCst);
assert!(
total_started < 50,
"Pending tasks must not have started after cancellation (started: {})",
total_started
);
}
#[test]
fn test_inspection_engine_progress_and_cancellation_api() {
let options = Options::default();
let engine = InspectionEngine::new(options);
assert_eq!(engine.completion_percentage(), 0.0);
assert!(!engine.is_cancelled());
let prog = engine.progress();
prog.set_percentage(75);
assert_eq!(engine.completion_percentage(), 75.0);
engine.cancel();
assert!(engine.is_cancelled());
assert!(engine.progress().is_cancelled());
}
}