use proofman_common::{ProofCtx, SetupCtx};
use proofman_fields::PrimeField64;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Instant;
use zisk_common::{stats_begin, stats_end, BusDevice, Instance, InstanceType, Stats};
use zisk_pil::{MainTraceRow, MainTraceRowPackedIndexed};
use zisk_sm_main::{MainInstance, MainPlanner, MainSmError};
use crate::error::{ExecutorError, ExecutorResult, RwLockExt};
use crate::state::ExecutionState;
pub struct WitnessGenerator {
chunk_size: u64,
packed: AtomicBool,
}
impl WitnessGenerator {
pub fn new(chunk_size: u64) -> Self {
Self { chunk_size, packed: AtomicBool::new(false) }
}
pub fn compute_main_witness<F: PrimeField64>(
&self,
pctx: &ProofCtx<F>,
state: &ExecutionState<F>,
main_instance: &MainInstance<F>,
trace_buffer: Vec<F>,
_caller_stats_id: u64,
) -> ExecutorResult<()> {
let witness_start_time = Instant::now();
let (airgroup_id, air_id) = pctx.dctx_get_instance_info(main_instance.ictx.global_id)?;
stats_begin!(state.stats, _caller_stats_id, _stats_scope, "AIR_MAIN_WITNESS", air_id);
let zisk_rom = state.get_rom()?;
let segment_id =
main_instance.ictx.plan.segment_id.ok_or(MainSmError::MissingSegmentId)?.as_usize();
let num_within = MainPlanner::traces_per_segment(self.chunk_size)?;
let (segment_min_traces, prev_chunk_last_c) = {
let min_traces_guard = state.min_traces.read_or_poison("min_traces")?;
let store = min_traces_guard.as_ref().ok_or(ExecutorError::MinTracesNotSet)?;
let start = segment_id * num_within;
let end = (start + num_within).min(store.len());
(
store.get(start..end).unwrap_or_default().to_vec(),
start.checked_sub(1).and_then(|i| store.get(i)).map(|t| t.last_c),
)
};
let air_instance = if self.packed.load(Ordering::Relaxed) {
main_instance.compute_witness::<MainTraceRowPackedIndexed<F>>(
&zisk_rom,
&segment_min_traces,
prev_chunk_last_c,
self.chunk_size,
trace_buffer,
)?
} else {
main_instance.compute_witness::<MainTraceRow<F>>(
&zisk_rom,
&segment_min_traces,
prev_chunk_last_c,
self.chunk_size,
trace_buffer,
)?
};
pctx.add_air_instance(air_instance, main_instance.ictx.global_id);
stats_end!(state.stats, &_stats_scope);
let stats = Stats::new_main_completed(airgroup_id, air_id, witness_start_time);
state.stats.insert_witness_stats(main_instance.ictx.global_id, stats);
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn compute_secn_witness<F: PrimeField64>(
&self,
pctx: &ProofCtx<F>,
sctx: &SetupCtx<F>,
state: &ExecutionState<F>,
global_id: usize,
secn_instance: &dyn Instance<F>,
collectors: Vec<(usize, Box<dyn BusDevice<u64>>)>,
trace_buffer: Vec<F>,
_caller_stats_id: u64,
) -> ExecutorResult<()> {
let witness_start_time = Instant::now();
let _stats_msg = match secn_instance.instance_type() {
InstanceType::Instance => "AIR_SECN_WITNESS",
InstanceType::Table => "AIR_WITNESS_TABLE",
};
let (_airgroup_id, _air_id) = pctx.dctx_get_instance_info(global_id)?;
stats_begin!(state.stats, _caller_stats_id, _stats_scope, _stats_msg, _air_id);
let air_instance = secn_instance.compute_witness(
pctx,
sctx,
collectors,
trace_buffer,
self.packed.load(Ordering::Relaxed),
)?;
if let Some(air_instance) = air_instance {
let should_add_instance = secn_instance.instance_type() == InstanceType::Instance
|| (secn_instance.instance_type() == InstanceType::Table
&& pctx.dctx_is_my_process_instance(global_id)?);
if should_add_instance {
pctx.add_air_instance(air_instance, global_id);
}
}
stats_end!(state.stats, &_stats_scope);
state.stats.set_witness_duration(global_id, witness_start_time.elapsed().as_millis());
Ok(())
}
pub fn set_packed(&self, packed: bool) {
self.packed.store(packed, Ordering::SeqCst);
}
pub fn is_packed(&self) -> bool {
self.packed.load(Ordering::Relaxed)
}
}