pub mod air_classifier;
pub mod collector;
pub mod generator;
pub mod handlers;
pub use air_classifier::*;
pub use collector::*;
pub use generator::*;
pub use handlers::*;
use std::collections::hash_map::Entry;
use std::collections::{BTreeMap, HashMap};
use std::sync::{Arc, Mutex};
use proofman_common::{BufferPool, ProofCtx, SetupCtx};
use proofman_fields::PrimeField64;
use zisk_asm_runner::AsmRunnerRH;
use zisk_common::{CheckPoint, InstanceCtx, InstanceType, Plan, StatsScope};
use zisk_core::ZiskRom;
use zisk_pil::RomTrace;
use zisk_sm_main::MainInstance;
use crate::error::{ExecutorError, ExecutorResult, MutexExt, RwLockExt};
use crate::ports::{Dctx, GlobalId, ProofRegistry};
use crate::sm::StaticSMBundle;
use crate::state::ExecutionState;
pub struct WitnessContext<'a, F: PrimeField64> {
pub pctx: &'a ProofCtx<F>,
pub sctx: &'a SetupCtx<F>,
pub state: &'a ExecutionState<F>,
pub buffer_pool: &'a dyn BufferPool<F>,
pub stats_scope: &'a StatsScope,
pub registry: &'a dyn Dctx,
pub is_asm_emulator: bool,
}
impl<'a, F: PrimeField64> WitnessContext<'a, F> {
#[allow(clippy::too_many_arguments)]
pub fn new(
pctx: &'a ProofCtx<F>,
sctx: &'a SetupCtx<F>,
state: &'a ExecutionState<F>,
buffer_pool: &'a dyn BufferPool<F>,
stats_scope: &'a StatsScope,
registry: &'a dyn Dctx,
is_asm_emulator: bool,
) -> Self {
Self { pctx, sctx, state, buffer_pool, stats_scope, registry, is_asm_emulator }
}
pub fn get_instance_info(&self, global_id: usize) -> ExecutorResult<(usize, usize)> {
let info = self.registry.instance_info(GlobalId(global_id))?;
Ok((info.airgroup_id, info.air_id))
}
}
pub struct WitnessPhase<F: PrimeField64> {
sm_bundle: Arc<StaticSMBundle<F>>,
collector: ChunkDataCollector<F>,
witness_generator: WitnessGenerator,
trace_buffer_rom: Mutex<Vec<F>>,
}
impl<F: PrimeField64> WitnessPhase<F> {
pub fn new(chunk_size: u64, sm_bundle: Arc<StaticSMBundle<F>>) -> Self {
let collector = ChunkDataCollector::new(sm_bundle.clone());
let witness_generator = WitnessGenerator::new(chunk_size);
let trace_buffer_rom = Mutex::new(vec![F::ZERO; RomTrace::<F>::NUM_ROWS]);
Self { sm_bundle, collector, witness_generator, trace_buffer_rom }
}
pub fn set_rh_data(&self, rh_data: AsmRunnerRH) -> ExecutorResult<()> {
self.collector.set_rh_data(rh_data)
}
pub fn set_rom(&self, zisk_rom: Arc<ZiskRom>) -> ExecutorResult<()> {
self.collector.set_rom(zisk_rom.clone())
}
pub fn set_packed(&self, packed: bool) {
self.witness_generator.set_packed(packed);
}
pub fn is_packed(&self) -> bool {
self.witness_generator.is_packed()
}
pub fn reset(&self) -> ExecutorResult<()> {
*self.trace_buffer_rom.lock_or_poison("trace_buffer_rom")? =
vec![F::ZERO; RomTrace::<F>::NUM_ROWS];
Ok(())
}
pub fn populate_main_instances(
&self,
registry: &dyn ProofRegistry,
state: &ExecutionState<F>,
assignments: Vec<(usize, Plan)>,
) -> ExecutorResult<()> {
let mut main_instances =
state.instance_set.main_instances.write_or_poison("main_instances")?;
for (global_id, plan) in assignments {
main_instances.entry(global_id).or_insert_with(|| {
std::sync::Arc::new(MainInstance::new(
InstanceCtx::new(global_id, plan),
self.sm_bundle.get_std(),
))
});
let gid = GlobalId(global_id);
if registry.is_my_process_instance(gid)? {
registry.set_witness_ready(gid, false);
}
}
Ok(())
}
pub fn configure_sm_instances(
&self,
pctx: &ProofCtx<F>,
plannings: &BTreeMap<usize, Vec<Plan>>,
) {
self.sm_bundle.configure_instances(pctx, plannings);
}
pub fn populate_secn_instances(
&self,
state: &ExecutionState<F>,
plans: Vec<Plan>,
) -> ExecutorResult<()> {
let mut secn_instances =
state.instance_set.secn_instances.write_or_poison("secn_instances")?;
for plan in plans {
let global_id =
plan.global_id.ok_or(ExecutorError::SecnPlanMissing { phase: "populate" })?;
if let Entry::Vacant(e) = secn_instances.entry(global_id) {
let instance = self.sm_bundle.build_instance(InstanceCtx::new(global_id, plan))?;
e.insert(instance);
}
}
Ok(())
}
pub fn configure_checkpoints(
&self,
registry: &dyn ProofRegistry,
state: &ExecutionState<F>,
global_ids: &[usize],
) -> ExecutorResult<()> {
let secn_instances = state.instance_set.secn_instances.read_or_poison("secn_instances")?;
for &global_id in global_ids {
let instance = secn_instances
.get(&global_id)
.ok_or(ExecutorError::InstanceNotFound { global_id })?;
instance.reset();
if instance.instance_type() == InstanceType::Instance {
let chunks: Vec<usize> = match instance.check_point() {
CheckPoint::None => vec![],
CheckPoint::Single(chunk_id) => vec![chunk_id.as_usize()],
CheckPoint::Multiple(chunk_ids) => {
chunk_ids.iter().map(|id| id.as_usize()).collect()
}
};
let gid = GlobalId(global_id);
let info = registry.instance_info(gid)?;
let is_memory_related = AirClassifier::is_memory_related(info.air_id);
registry.set_chunks(gid, &chunks, is_memory_related);
}
}
Ok(())
}
pub fn dispatch(&self, ctx: &WitnessContext<'_, F>, global_id: usize) -> ExecutorResult<()> {
let (airgroup_id, air_id) = ctx.get_instance_info(global_id)?;
let stats_scope_id = ctx.stats_scope.id();
if AirClassifier::is_main(air_id) {
return MainWitnessHandler::dispatch(
&self.witness_generator,
ctx.state,
ctx.pctx,
global_id,
ctx.buffer_pool,
stats_scope_id,
);
}
let instance_type = {
let secn = ctx.state.instance_set.secn_instances.read_or_poison("secn_instances")?;
secn.get(&global_id)
.ok_or(ExecutorError::InstanceNotFound { global_id })?
.instance_type()
};
match instance_type {
InstanceType::Table => TableWitnessHandler::dispatch(
&self.witness_generator,
ctx.state,
ctx.pctx,
ctx.sctx,
global_id,
ctx.buffer_pool,
stats_scope_id,
),
InstanceType::Instance => {
if AirClassifier::is_rom(airgroup_id, air_id) {
self.rom_dispatch(ctx, global_id, airgroup_id, air_id, stats_scope_id)
} else {
SecondaryWitnessHandler::dispatch(
&self.witness_generator,
&self.collector,
ctx.state,
ctx.pctx,
ctx.sctx,
global_id,
ctx.buffer_pool,
stats_scope_id,
)
}
}
}
}
fn rom_dispatch(
&self,
ctx: &WitnessContext<'_, F>,
global_id: usize,
airgroup_id: usize,
air_id: usize,
stats_scope_id: u64,
) -> ExecutorResult<()> {
let secn_instances =
ctx.state.instance_set.secn_instances.read_or_poison("secn_instances")?;
let secn_instance =
secn_instances.get(&global_id).ok_or(ExecutorError::InstanceNotFound { global_id })?;
let needs_collection = !ctx
.state
.collector_store
.inner
.read_or_poison("collector_store")?
.contains_key(&global_id);
let instance = &**secn_instance;
if needs_collection {
if ctx.is_asm_emulator {
ctx.state.register_empty_collector(global_id, airgroup_id, air_id)?;
} else {
self.collector.collect_single(ctx.pctx, ctx.state, global_id, instance)?;
}
}
let collectors =
ctx.state.take_collectors_for_instance(global_id, instance.instance_type())?;
let trace_buffer =
std::mem::take(&mut *self.trace_buffer_rom.lock_or_poison("trace_buffer_rom")?);
self.witness_generator.compute_secn_witness(
ctx.pctx,
ctx.sctx,
ctx.state,
global_id,
instance,
collectors,
trace_buffer,
stats_scope_id,
)
}
pub fn pre_calculate(
&self,
pctx: &ProofCtx<F>,
registry: &dyn Dctx,
state: &ExecutionState<F>,
global_ids: &[usize],
is_asm_emulator: bool,
) -> ExecutorResult<()> {
let secn_instances_guard =
state.instance_set.secn_instances.read_or_poison("secn_instances")?;
let mut instances_to_collect = HashMap::new();
for &global_id in global_ids {
let info = registry.instance_info(GlobalId(global_id))?;
if AirClassifier::is_main(info.air_id) {
registry.set_witness_ready(GlobalId(global_id), false);
} else if AirClassifier::is_rom(info.airgroup_id, info.air_id) {
if is_asm_emulator {
registry.set_witness_ready(GlobalId(global_id), false);
} else {
handlers::rom_rust::pre_calculate(
registry,
state,
&secn_instances_guard,
&mut instances_to_collect,
global_id,
info.airgroup_id,
info.air_id,
)?;
}
} else {
self.handle_secondary_pre_calculate(
registry,
state,
&secn_instances_guard,
&mut instances_to_collect,
global_id,
)?;
}
}
if !instances_to_collect.is_empty() {
self.collector.collect(pctx, state, instances_to_collect)?;
}
Ok(())
}
fn handle_secondary_pre_calculate<'a>(
&self,
registry: &dyn Dctx,
state: &ExecutionState<F>,
secn_instances: &'a SecnInstanceMap<F>,
instances_to_collect: &mut SecnInstanceMapRef<'a, F>,
global_id: usize,
) -> ExecutorResult<()> {
let secn_instance =
secn_instances.get(&global_id).ok_or(ExecutorError::InstanceNotFound { global_id })?;
if secn_instance.instance_type() == InstanceType::Instance
&& !state
.collector_store
.inner
.read_or_poison("collector_store")?
.contains_key(&global_id)
{
instances_to_collect.insert(global_id, &**secn_instance);
} else {
registry.set_witness_ready(GlobalId(global_id), true);
}
Ok(())
}
}