use std::sync::{Arc, Mutex};
use crate::bus::pub_outs_collector::PubOutsCollector;
use crate::error::{ExecutorError, ExecutorResult, MutexExt};
use crate::execution::output::{BackendArtifacts, ExecutionOutput};
use crate::{CountersChunkMetrics, MAX_NUM_STEPS};
use super::{AsmResources, AsmRunnerSupervisor, AsmTransport, MtChunkProcessor};
use proofman_fields::PrimeField64;
use zisk_asm_runner::{AsmRunnerMT, HintsShmem};
use zisk_common::{
io::StreamSource, io::ZiskStdin, stats_begin, stats_end, AsmExecutionInfo, ChunkId, EmuTrace,
ExecutorStatsHandle, StatsScope,
};
use zisk_core::ZiskRom;
use zisk_precomp_hints::HintsProcessor;
pub struct EmulatorAsm {
chunk_size: u64,
transport: AsmTransport,
asm_execution_info: Mutex<Option<AsmExecutionInfo>>,
}
impl EmulatorAsm {
pub fn new(chunk_size: u64) -> Self {
Self { chunk_size, transport: AsmTransport::new(), asm_execution_info: Mutex::new(None) }
}
pub fn get_asm_execution_info(&self) -> ExecutorResult<Option<AsmExecutionInfo>> {
Ok(self.asm_execution_info.lock_or_poison("asm_execution_info")?.clone())
}
pub fn set_asm_resources(&self, asm_resources: Arc<AsmResources>) -> ExecutorResult<()> {
self.transport.set_asm_resources(asm_resources)
}
pub fn reset(&self) -> ExecutorResult<()> {
self.transport.reset()
}
pub fn signal_cancellation(&self) -> ExecutorResult<()> {
self.transport.signal_cancellation()
}
pub fn get_hints_processor(&self) -> ExecutorResult<Arc<HintsProcessor<HintsShmem>>> {
self.transport.get_hints_processor()
}
pub fn set_active_services(&self, is_first_process: bool) -> ExecutorResult<()> {
self.transport.set_active_services(is_first_process)
}
pub fn set_hints_stream_src(&self, stream: StreamSource) -> ExecutorResult<()> {
self.transport.set_hints_stream_src(stream)
}
pub fn set_inputs_stream_src(&self, stream: StreamSource) -> ExecutorResult<()> {
self.transport.set_inputs_stream_src(stream)
}
pub fn submit_hint_direct(&self, data: &[u64]) -> ExecutorResult<()> {
self.transport.submit_hint_direct(data)
}
pub fn append_raw_input(&self, bytes: &[u8]) -> ExecutorResult<()> {
self.transport.append_raw_input(bytes)
}
#[allow(clippy::too_many_arguments)]
pub fn execute<F: PrimeField64>(
&self,
zisk_rom: &ZiskRom,
stdin: &ZiskStdin,
has_rom_sm: bool,
use_hints: bool,
stats: &ExecutorStatsHandle,
_caller_stats_scope: &StatsScope,
chunk_hook: crate::ChunkHook<'_>,
) -> ExecutorResult<ExecutionOutput> {
let asm_resources = self.transport.resources()?;
let has_hints_stream = asm_resources.is_hints_stream_initialized();
if use_hints && has_hints_stream {
asm_resources.start_stream()?;
}
if asm_resources.is_inputs_stream_initialized() {
asm_resources.start_inputs_stream()?;
}
stats_begin!(stats, _caller_stats_scope, _exec_scope, "EXECUTE_WITH_ASSEMBLY", 0);
stats_begin!(stats, &_exec_scope, _write_scope, "ASM_WRITE_INPUT", 0);
asm_resources.write_input(stdin)?;
stats_end!(stats, &_write_scope);
let supervisor =
AsmRunnerSupervisor::spawn_on(&asm_resources, self.chunk_size, has_rom_sm, stats);
let mt_result = self.run_mt_assembly::<F>(zisk_rom, stats, chunk_hook);
let output = match mt_result {
Ok((min_traces, counters, pub_outs)) => {
let steps = min_traces.iter().map(|trace| trace.steps).sum::<u64>();
let (handle_mo, handle_rh) = supervisor.into_handles();
Ok(ExecutionOutput {
min_traces,
counters,
pub_outs,
steps,
backend: BackendArtifacts::Asm { mo: Some(handle_mo), rh: handle_rh },
})
}
Err(e) => {
supervisor.cleanup_after_mt_failure(|| asm_resources.signal_cancellation());
Err(e)
}
};
stats_end!(stats, &_exec_scope);
output
}
fn run_mt_assembly<F: PrimeField64>(
&self,
zisk_rom: &ZiskRom,
stats: &ExecutorStatsHandle,
chunk_hook: crate::ChunkHook<'_>,
) -> ExecutorResult<(Vec<std::sync::Arc<EmuTrace>>, CountersChunkMetrics, PubOutsCollector)>
{
stats_begin!(stats, 0, _mt_scope, "RUN_MT_ASSEMBLY", 0);
let processor: MtChunkProcessor<F> = MtChunkProcessor::new();
#[allow(unused_variables)]
let mt_scope_id = _mt_scope.id();
let scope_result: ExecutorResult<_> = rayon::in_place_scope(|scope| {
let processor_ref = &processor;
let on_chunk = |idx: usize,
emu_traces: &[std::sync::Arc<EmuTrace>],
last: bool|
-> anyhow::Result<()> {
let emu_trace = emu_traces[idx].clone();
let chunk_id = ChunkId(idx);
scope.spawn(move |_| {
processor_ref.process_chunk(chunk_id, &emu_trace, zisk_rom, stats, mt_scope_id);
});
chunk_hook(idx, emu_traces, last).map_err(anyhow::Error::from)
};
let asm_resources = self.transport.resources()?;
let mt_shmem = &mut asm_resources.readers().mt.lock_or_poison("mt_shmem_reader")?;
let asm_resources_for_failure = asm_resources.clone();
let result = AsmRunnerMT::run_and_count(
mt_shmem,
MAX_NUM_STEPS,
self.chunk_size,
on_chunk,
move || {
asm_resources_for_failure.signal_cancellation().map_err(anyhow::Error::from)
},
asm_resources.asm_services().clone(),
stats.clone(),
)
.map_err(ExecutorError::asm_backend)?;
Ok(result)
});
let (emu_traces, asm_execution_info) = scope_result?;
self.asm_execution_info.lock_or_poison("asm_execution_info")?.replace(asm_execution_info);
let (counters, pub_outs) = processor.finalize()?;
stats_end!(stats, &_mt_scope);
Ok((emu_traces, counters, pub_outs))
}
}