use crossbeam::atomic::AtomicCell;
use proofman_common::ProofCtx;
use proofman_fields::PrimeField64;
use rayon::prelude::*;
use std::{
collections::HashMap,
sync::{
atomic::{AtomicUsize, Ordering},
Arc, Mutex, RwLock,
},
time::Instant,
};
use tracing::error;
use zisk_common::{
CheckPoint, ChunkId, DataBusTrait, EmuTrace, ExecutorStatsHandle, Instance, PayloadType, Stats,
};
use zisk_core::ZiskRom;
use ziskemu::ZiskEmulator;
use crate::error::{ExecutorError, ExecutorResult, RwLockExt};
use crate::{state::ChunkCollector, ExecutionState, StaticDataBusCollect, StaticSMBundle};
use zisk_asm_runner::AsmRunnerRH;
type CollectorSlots = Arc<RwLock<HashMap<usize, Vec<Option<ChunkCollector>>>>>;
struct WorkerCtx<'a, F: PrimeField64> {
next_chunk: &'a AtomicUsize,
ordered_chunks: &'a [usize],
chunks_to_execute: &'a [Vec<usize>],
data_buses: &'a [Mutex<Option<StaticDataBusCollect<PayloadType, F>>>],
zisk_rom: &'a ZiskRom,
min_traces: &'a [Arc<EmuTrace>],
pctx: &'a ProofCtx<F>,
collectors_by_instance: &'a CollectorSlots,
n_chunks_left: &'a [AtomicUsize],
collect_start_times: &'a [AtomicCell<Option<Instant>>],
stats: &'a ExecutorStatsHandle,
global_ids_map: &'a HashMap<usize, usize>,
global_id_chunks: &'a HashMap<usize, Vec<usize>>,
errors: &'a Mutex<Vec<String>>,
}
#[inline]
fn push_error(errors: &Mutex<Vec<String>>, message: String) {
if let Ok(mut errs) = errors.lock() {
errs.push(message);
}
}
pub struct ChunkDataCollector<F: PrimeField64> {
sm_bundle: Arc<StaticSMBundle<F>>,
}
impl<F: PrimeField64> ChunkDataCollector<F> {
pub fn new(sm_bundle: Arc<StaticSMBundle<F>>) -> Self {
Self { sm_bundle }
}
pub fn set_rom(&self, zisk_rom: Arc<ZiskRom>) -> ExecutorResult<()> {
self.sm_bundle.set_rom(zisk_rom)
}
pub fn set_rh_data(&self, rh_data: AsmRunnerRH) -> ExecutorResult<()> {
self.sm_bundle.set_rh_data(rh_data)
}
pub fn compute_chunks_to_execute(
&self,
min_traces: &[Arc<EmuTrace>],
secn_instances: &HashMap<usize, &dyn Instance<F>>,
) -> (Vec<Vec<usize>>, HashMap<usize, Vec<usize>>) {
let mut chunks_to_execute = vec![Vec::new(); min_traces.len()];
let mut global_id_chunks: HashMap<usize, Vec<usize>> = HashMap::new();
secn_instances.iter().for_each(|(global_idx, secn_instance)| {
match secn_instance.check_point() {
CheckPoint::None => {}
CheckPoint::Single(chunk_id) => {
chunks_to_execute[chunk_id.as_usize()].push(*global_idx);
global_id_chunks.entry(*global_idx).or_default().push(chunk_id.as_usize());
}
CheckPoint::Multiple(chunk_ids) => {
chunk_ids.iter().for_each(|&chunk_id| {
chunks_to_execute[chunk_id.as_usize()].push(*global_idx);
global_id_chunks.entry(*global_idx).or_default().push(chunk_id.as_usize());
});
}
}
});
for chunk_ids in global_id_chunks.values_mut() {
chunk_ids.sort();
}
(chunks_to_execute, global_id_chunks)
}
pub fn order_chunks(
&self,
chunks_to_execute: &[Vec<usize>],
global_id_chunks: &HashMap<usize, Vec<usize>>,
) -> Vec<usize> {
let mut ordered_chunks = Vec::new();
let mut already_selected_chunks = vec![false; chunks_to_execute.len()];
let mut n_global_ids_incompleted = global_id_chunks.len();
let mut n_chunks_by_global_id: HashMap<usize, usize> =
global_id_chunks.iter().map(|(global_id, chunks)| (*global_id, chunks.len())).collect();
while n_global_ids_incompleted > 0 {
let selected_global_id = n_chunks_by_global_id
.iter()
.filter(|(_, &count)| count > 0)
.min_by_key(|(_, &count)| count)
.map(|(&global_id, _)| global_id);
if let Some(global_id) = selected_global_id {
for chunk_id in global_id_chunks[&global_id].iter() {
if already_selected_chunks[*chunk_id] {
continue;
}
ordered_chunks.push(*chunk_id);
already_selected_chunks[*chunk_id] = true;
for global_idx in chunks_to_execute[*chunk_id].iter() {
if let Some(count) = n_chunks_by_global_id.get_mut(global_idx) {
*count -= 1;
if *count == 0 {
n_chunks_by_global_id.remove(global_idx);
n_global_ids_incompleted -= 1;
}
}
}
}
} else {
break;
}
}
ordered_chunks
}
pub fn collect_single(
&self,
pctx: &ProofCtx<F>,
state: &ExecutionState<F>,
global_id: usize,
instance: &dyn Instance<F>,
) -> ExecutorResult<()> {
let mut map = HashMap::with_capacity(1);
map.insert(global_id, instance);
self.collect(pctx, state, map)?;
Ok(())
}
pub fn collect(
&self,
pctx: &ProofCtx<F>,
state: &ExecutionState<F>,
secn_instances: HashMap<usize, &dyn Instance<F>>,
) -> ExecutorResult<()> {
let min_traces_guard = state.min_traces.read_or_poison("min_traces")?;
let min_traces = min_traces_guard.as_ref().ok_or(ExecutorError::MinTracesNotSet)?;
let (chunks_to_execute, global_id_chunks) =
self.compute_chunks_to_execute(min_traces, &secn_instances);
let ordered_chunks = self.order_chunks(&chunks_to_execute, &global_id_chunks);
let global_ids: Vec<usize> = secn_instances.keys().copied().collect();
let collect_start_times: Vec<AtomicCell<Option<Instant>>> =
global_ids.iter().map(|_| AtomicCell::new(None)).collect();
let global_ids_map: HashMap<usize, usize> =
global_ids.iter().enumerate().map(|(idx, &id)| (id, idx)).collect();
let zisk_rom = state.get_rom()?;
let data_buses: Vec<Option<_>> = chunks_to_execute
.par_iter()
.enumerate()
.map(|(chunk_id, global_idxs)| {
if global_idxs.is_empty() {
Ok(None)
} else {
crate::StaticDataBusCollect::for_chunk(
pctx,
&secn_instances,
ChunkId(chunk_id),
global_idxs,
zisk_rom.as_ref(),
)
.map(Some)
}
})
.collect::<ExecutorResult<_>>()?;
let data_buses: Vec<_> = data_buses.into_iter().map(Mutex::new).collect();
let n_chunks_left: Vec<AtomicUsize> = global_ids
.iter()
.map(|global_id| {
global_id_chunks.get(global_id).map(|chunks| AtomicUsize::new(chunks.len())).ok_or(
ExecutorError::MissingIndexEntry {
global_id: *global_id,
index: "global_id_chunks",
},
)
})
.collect::<ExecutorResult<Vec<_>>>()?;
for global_id in global_ids.iter() {
let (airgroup_id, air_id) = pctx.dctx_get_instance_info(*global_id)?;
let n_chunks = global_id_chunks
.get(global_id)
.ok_or(ExecutorError::MissingIndexEntry {
global_id: *global_id,
index: "global_id_chunks",
})?
.len();
let stats = Stats::new_pending_collection(airgroup_id, air_id, n_chunks);
state
.collector_store
.inner
.write_or_poison("collector_store")?
.insert(*global_id, (0..n_chunks).map(|_| None).collect());
state.stats.insert_witness_stats(*global_id, stats);
}
let next_chunk = AtomicUsize::new(0);
let zisk_rom = state.get_rom()?;
let errors: Mutex<Vec<String>> = Mutex::new(Vec::new());
let ctx = WorkerCtx {
next_chunk: &next_chunk,
ordered_chunks: &ordered_chunks,
chunks_to_execute: &chunks_to_execute,
data_buses: &data_buses,
zisk_rom: &zisk_rom,
min_traces,
pctx,
collectors_by_instance: &state.collector_store.inner,
n_chunks_left: &n_chunks_left,
collect_start_times: &collect_start_times,
stats: &state.stats,
global_ids_map: &global_ids_map,
global_id_chunks: &global_id_chunks,
errors: &errors,
};
rayon::in_place_scope(|scope| {
for _ in 0..rayon::current_num_threads() {
let ctx = &ctx;
scope.spawn(move |_| Self::worker_loop(ctx));
}
});
let err_vec = errors.lock().unwrap_or_else(|poisoned| {
error!("errors mutex was poisoned during parallel chunk execution");
poisoned.into_inner()
});
if !err_vec.is_empty() {
let message = err_vec
.iter()
.enumerate()
.map(|(i, e)| format!("[Error {}] {e}", i + 1))
.collect::<Vec<_>>()
.join("\n");
return Err(ExecutorError::MtChunkProcessing { count: err_vec.len(), message });
}
Ok(())
}
fn worker_loop(ctx: &WorkerCtx<'_, F>) {
loop {
let next_chunk_id = ctx.next_chunk.fetch_add(1, Ordering::Relaxed);
if next_chunk_id >= ctx.ordered_chunks.len() {
break;
}
let chunk_id = ctx.ordered_chunks[next_chunk_id];
let data_bus = match ctx.data_buses[chunk_id].lock() {
Ok(mut lock) => match lock.take() {
Some(bus) => bus,
None => continue,
},
Err(_) => {
push_error(
ctx.errors,
format!("data_buses lock poisoned for chunk {chunk_id}"),
);
continue;
}
};
Self::process_one_chunk(chunk_id, data_bus, ctx);
}
}
fn process_one_chunk(
chunk_id: usize,
mut data_bus: StaticDataBusCollect<PayloadType, F>,
ctx: &WorkerCtx<'_, F>,
) {
let mut affected_globals: Vec<(usize, usize)> = Vec::new();
for global_id in ctx.chunks_to_execute[chunk_id].iter() {
match ctx.global_ids_map.get(global_id) {
Some(&global_id_idx) => {
let start_time_cell = &ctx.collect_start_times[global_id_idx];
if start_time_cell.load().is_none() {
start_time_cell.store(Some(Instant::now()));
}
affected_globals.push((*global_id, global_id_idx));
}
None => {
push_error(ctx.errors, format!("global_id {global_id} not in global_ids_map"));
}
}
}
ZiskEmulator::process_emu_traces::<F, _, _>(
ctx.zisk_rom,
&ctx.min_traces[chunk_id],
&mut data_bus,
);
let devices = data_bus.into_devices(false);
let mut entries: Vec<(usize, usize, Option<ChunkCollector>)> = Vec::new();
for (global_id, col) in devices {
match ctx.global_id_chunks.get(&global_id) {
Some(chunk_order) => {
if let Some(position) = chunk_order.iter().position(|&id| id == chunk_id) {
entries.push((global_id, position, Some((chunk_id, col))));
} else {
push_error(
ctx.errors,
format!(
"chunk_id {chunk_id} not in chunk_order for global_id {global_id}"
),
);
}
}
None => {
push_error(
ctx.errors,
format!("global_id {global_id} not found in global_id_chunks"),
);
}
}
}
match ctx.collectors_by_instance.write() {
Ok(mut guard) => {
for (global_id, position, entry) in entries {
if let Some(vec) = guard.get_mut(&global_id) {
vec[position] = entry;
} else {
push_error(
ctx.errors,
format!("global_id {global_id} not in collectors_by_instance"),
);
}
}
}
Err(_) => {
push_error(ctx.errors, "collectors_by_instance lock poisoned".to_string());
}
}
for (global_id, global_id_idx) in affected_globals {
if ctx.n_chunks_left[global_id_idx].fetch_sub(1, Ordering::SeqCst) == 1 {
ctx.pctx.set_witness_ready(global_id, true);
Self::record_completion_stats(global_id, global_id_idx, ctx);
}
}
}
fn record_completion_stats(global_id: usize, global_id_idx: usize, ctx: &WorkerCtx<'_, F>) {
let Some(collect_start_time) = ctx.collect_start_times[global_id_idx].load() else {
push_error(ctx.errors, format!("collect_start_time not set for global_id {global_id}"));
return;
};
let collect_duration = collect_start_time.elapsed().as_millis() as u64;
match (ctx.pctx.dctx_get_instance_info(global_id), ctx.global_id_chunks.get(&global_id)) {
(Ok((airgroup_id, air_id)), Some(chunks)) => {
let new_stats = Stats::new_with_collection(
airgroup_id,
air_id,
chunks.len(),
collect_start_time,
collect_duration,
);
ctx.stats.insert_witness_stats(global_id, new_stats);
}
(Err(e), _) => {
push_error(
ctx.errors,
format!("Failed to get instance info for global_id {global_id}: {e}"),
);
}
(Ok(_), None) => {
push_error(ctx.errors, format!("global_id {global_id} not in global_id_chunks"));
}
}
}
}