use std::path::Path;
use std::sync::{Arc, Mutex};
use zisk_asm_runner::{
AsmRunnerOptions, AsmServices, ControlShmem, GpuBufferSource, HintsShmem, InputsShmemWriter,
};
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
use zisk_asm_runner::{MOShmemReader, MTShmemReader, RHShmemReader};
use zisk_common::io::{StreamSink, StreamSource, ZiskStdin, ZiskStream};
use zisk_precomp_hints::{HintsProcessor, MpiBroadcastFn};
use crate::error::{ExecutorError, ExecutorResult, MutexExt};
#[derive(Clone)]
pub struct AsmResourcesConfig {
pub local_rank: i32,
pub unlock_mapped_memory: bool,
}
impl std::fmt::Debug for AsmResourcesConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AsmResourcesConfig")
.field("local_rank", &self.local_rank)
.field("unlock_mapped_memory", &self.unlock_mapped_memory)
.finish_non_exhaustive()
}
}
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
pub struct AsmShmemReaders {
pub mt: Arc<Mutex<MTShmemReader>>,
pub mo: Arc<Mutex<MOShmemReader>>,
pub rh: Arc<Mutex<Option<RHShmemReader>>>,
}
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
impl AsmShmemReaders {
fn new(
shm_prefix: &str,
unlock_mapped_memory: bool,
gpu_buffer_src: GpuBufferSource,
) -> ExecutorResult<Self> {
Ok(Self {
mt: Arc::new(Mutex::new(
MTShmemReader::new(shm_prefix, unlock_mapped_memory)
.map_err(ExecutorError::asm_backend)?,
)),
mo: Arc::new(Mutex::new(
MOShmemReader::new(shm_prefix, unlock_mapped_memory, gpu_buffer_src)
.map_err(ExecutorError::asm_backend)?,
)),
rh: Arc::new(Mutex::new(None)),
})
}
}
pub struct AsmSharedResources {
config: AsmResourcesConfig,
pub shmem_inputs: Arc<InputsShmemWriter>,
hints_stream: Option<Arc<Mutex<ZiskStream<HintsProcessor<HintsShmem>>>>>,
inputs_stream: Arc<Mutex<ZiskStream<InputsShmemWriter>>>,
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
pub readers: AsmShmemReaders,
}
impl std::fmt::Debug for AsmSharedResources {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let hints_init =
self.hints_stream.as_ref().and_then(|s| s.lock().ok().map(|g| g.is_initialized()));
f.debug_struct("AsmSharedResources")
.field("config", &self.config)
.field("hints_stream", &hints_init)
.finish_non_exhaustive()
}
}
impl AsmSharedResources {
#[allow(clippy::too_many_arguments)]
pub fn new(
local_rank: i32,
unlock_mapped_memory: bool,
verbose_mode: proofman_common::VerboseMode,
mpi_broadcast_fn: Option<MpiBroadcastFn>,
init_rom: bool,
with_hints: bool,
shm_prefix: &str,
gpu_buffer_src: GpuBufferSource,
) -> ExecutorResult<Self> {
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
let readers = AsmShmemReaders::new(shm_prefix, unlock_mapped_memory, gpu_buffer_src)?;
#[cfg(not(all(target_os = "linux", target_arch = "x86_64")))]
let _ = gpu_buffer_src;
let control_writer = Arc::new(
ControlShmem::new(shm_prefix, unlock_mapped_memory)
.map_err(ExecutorError::asm_backend)?,
);
let config = AsmResourcesConfig { local_rank, unlock_mapped_memory };
let shmem_inputs = Arc::new(
InputsShmemWriter::new(shm_prefix, unlock_mapped_memory, control_writer.clone())
.map_err(ExecutorError::asm_backend)?,
);
let inputs_stream = Arc::new(Mutex::new(ZiskStream::from_arc(Arc::clone(&shmem_inputs))));
let hints_stream = if with_hints {
let active_services =
if init_rom { &AsmServices::SERVICES[..] } else { &AsmServices::SERVICES[..2] };
let hints_shmem = Arc::new(
HintsShmem::new(shm_prefix, unlock_mapped_memory, control_writer, active_services)
.map_err(ExecutorError::asm_backend)?,
);
let mut builder = HintsProcessor::builder(hints_shmem, Some(shmem_inputs.clone()))
.enable_stats(verbose_mode != proofman_common::VerboseMode::Info);
if let Some(broadcast_fn) = mpi_broadcast_fn {
builder = builder.with_mpi_broadcast(move |data| broadcast_fn(data));
}
let hints_processor = builder.build().map_err(ExecutorError::asm_backend)?;
Some(Arc::new(Mutex::new(ZiskStream::new(hints_processor))))
} else {
None
};
Ok(Self {
config,
hints_stream,
inputs_stream,
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
readers,
shmem_inputs,
})
}
}
pub struct AsmResources {
shared: Arc<AsmSharedResources>,
asm_services: AsmServices,
}
impl std::fmt::Debug for AsmResources {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AsmResources").field("shared", &self.shared).finish_non_exhaustive()
}
}
impl AsmResources {
pub fn new(shared: Arc<AsmSharedResources>, asm_services: AsmServices) -> ExecutorResult<Self> {
let sem_prefix = asm_services.sem_prefix();
shared.shmem_inputs.bind_semaphores(sem_prefix).map_err(ExecutorError::asm_backend)?;
if let Some(hints_stream) = &shared.hints_stream {
let processor = hints_stream.lock_or_poison("hints_stream")?.get_processor();
processor
.hints_sink()
.bind_semaphores(sem_prefix)
.map_err(ExecutorError::asm_backend)?;
}
Ok(Self { shared, asm_services })
}
pub fn new_standalone(
elf_hash: String,
asm_mt_path: &Path,
with_hints: bool,
verbose: proofman_common::VerboseMode,
gpu: bool,
) -> ExecutorResult<Self> {
let options = AsmRunnerOptions::new().with_local_rank(0);
let services = AsmServices::new(0, 0, elf_hash, asm_mt_path, with_hints, options)
.map_err(ExecutorError::asm_backend)?;
let gpu_buffer_src =
if gpu { GpuBufferSource::SelfAllocated } else { GpuBufferSource::Cpu };
let shared = Arc::new(AsmSharedResources::new(
0,
false,
verbose,
None,
true,
with_hints,
services.shm_prefix(),
gpu_buffer_src,
)?);
Self::new(shared, services)
}
pub fn get_hints_processor(&self) -> ExecutorResult<Arc<HintsProcessor<HintsShmem>>> {
self.shared
.hints_stream
.as_ref()
.ok_or(ExecutorError::HintsNotConfigured)?
.lock()
.map(|g| g.get_processor())
.map_err(|_| ExecutorError::mutex_poisoned("hints_stream"))
}
pub fn set_active_services(&self, is_first_process: bool) -> ExecutorResult<()> {
let Some(hints_stream) = &self.shared.hints_stream else { return Ok(()) };
let processor = hints_stream.lock_or_poison("hints_stream")?.get_processor();
let services =
if is_first_process { &AsmServices::SERVICES[..] } else { &AsmServices::SERVICES[..2] };
processor.hints_sink().set_active_services(services).map_err(ExecutorError::asm_backend)
}
pub fn submit_hint_direct(&self, data: &[u64]) -> ExecutorResult<()> {
let processor = self.get_hints_processor()?;
processor.hints_sink().submit(data).map_err(ExecutorError::asm_backend)
}
pub fn start_stream(&self) -> ExecutorResult<()> {
let Some(hints_stream) = &self.shared.hints_stream else {
return Err(ExecutorError::HintsNotConfigured);
};
hints_stream
.lock_or_poison("hints_stream")?
.start_stream()
.map_err(ExecutorError::asm_backend)
}
pub fn set_hints_stream_src(&self, stream: StreamSource) -> ExecutorResult<()> {
let Some(hints_stream) = &self.shared.hints_stream else {
return Err(ExecutorError::HintsNotConfigured);
};
hints_stream
.lock_or_poison("hints_stream")?
.set_stream_src(stream)
.map_err(ExecutorError::asm_backend)
}
pub fn is_hints_stream_initialized(&self) -> bool {
self.shared
.hints_stream
.as_ref()
.and_then(|s| s.lock().ok().map(|g| g.is_initialized()))
.unwrap_or(false)
}
pub fn signal_cancellation(&self) -> ExecutorResult<()> {
self.shared.shmem_inputs.signal_reset().map_err(ExecutorError::asm_backend)
}
pub fn reset(&self) {
if let Some(s) = &self.shared.hints_stream {
s.lock().expect("hints_stream mutex poisoned").reset();
}
self.shared.shmem_inputs.reset();
}
pub fn config(&self) -> &AsmResourcesConfig {
&self.shared.config
}
pub fn asm_services(&self) -> &AsmServices {
&self.asm_services
}
pub fn set_inputs_stream_src(&self, stream: StreamSource) -> ExecutorResult<()> {
self.shared
.inputs_stream
.lock_or_poison("inputs_stream")?
.set_stream_src(stream)
.map_err(ExecutorError::asm_backend)?;
Ok(())
}
pub fn start_inputs_stream(&self) -> ExecutorResult<()> {
self.shared
.inputs_stream
.lock_or_poison("inputs_stream")?
.start_stream()
.map_err(ExecutorError::asm_backend)
}
pub fn is_inputs_stream_initialized(&self) -> bool {
self.shared.inputs_stream.lock().map(|s| s.is_initialized()).unwrap_or(false)
}
pub fn write_input(&self, stdin: &ZiskStdin) -> ExecutorResult<()> {
stdin
.with_data(|data| self.shared.shmem_inputs.write_input(data))
.map_err(ExecutorError::asm_backend)
}
pub fn append_raw_input(&self, bytes: &[u8]) -> ExecutorResult<()> {
self.shared.shmem_inputs.append_input(bytes).map_err(ExecutorError::asm_backend)
}
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
pub fn readers(&self) -> &AsmShmemReaders {
&self.shared.readers
}
}
impl Drop for AsmResources {
fn drop(&mut self) {
self.shared.shmem_inputs.unbind_semaphores();
if let Some(hints_stream) = &self.shared.hints_stream {
if let Ok(g) = hints_stream.lock() {
g.get_processor().hints_sink().unbind_semaphores();
}
}
}
}