use std::{sync::Arc, thread::JoinHandle};
use zisk_asm_runner::{AsmRunnerMO, AsmRunnerRH};
use zisk_common::{EmuTrace, Plan};
use crate::error::{ExecutorError, ExecutorResult};
use crate::pub_outs_collector::PubOutsCollector;
use crate::CountersChunkMetrics;
pub struct ExecutionOutput {
pub min_traces: Vec<Arc<EmuTrace>>,
pub counters: CountersChunkMetrics,
pub pub_outs: PubOutsCollector,
pub steps: u64,
pub backend: BackendArtifacts,
}
pub enum BackendArtifacts {
Asm {
mo: Option<JoinHandle<ExecutorResult<AsmRunnerMO>>>,
rh: Option<JoinHandle<ExecutorResult<AsmRunnerRH>>>,
},
Rust,
}
impl BackendArtifacts {
pub fn await_mem_plans(&mut self) -> ExecutorResult<(Vec<Plan>, Option<u64>)> {
match self {
Self::Asm { mo, .. } => {
let handle = mo.take().ok_or(ExecutorError::RunnerHandleConsumed { name: "MO" })?;
let asm_runner_mo = handle
.join()
.map_err(|_| ExecutorError::RunnerThreadPanicked { name: "MO" })?
.map_err(|e| ExecutorError::RunnerFailed {
name: "MO",
message: e.to_string(),
})?;
Ok((asm_runner_mo.plans, asm_runner_mo.gpu_mops_used_bytes))
}
Self::Rust => Ok((Vec::new(), None)),
}
}
pub fn await_rom_histogram(&mut self) -> ExecutorResult<Option<AsmRunnerRH>> {
match self {
Self::Asm { rh, .. } => {
let Some(handle) = rh.take() else {
return Ok(None);
};
let rh_data = handle
.join()
.map_err(|_| ExecutorError::RunnerThreadPanicked { name: "RH" })?
.map_err(|e| ExecutorError::RunnerFailed {
name: "RH",
message: e.to_string(),
})?;
Ok(Some(rh_data))
}
Self::Rust => Ok(None),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use zisk_asm_runner::AsmRHData;
#[test]
fn rust_await_mem_plans_yields_empty() {
let mut backend = BackendArtifacts::Rust;
let (plans, gpu_mops_used_bytes) =
backend.await_mem_plans().expect("await_mem_plans on Rust");
assert!(plans.is_empty());
assert!(gpu_mops_used_bytes.is_none());
}
#[test]
fn rust_await_rom_histogram_yields_none() {
let mut backend = BackendArtifacts::Rust;
let rh = backend.await_rom_histogram().expect("await_rom_histogram on Rust");
assert!(rh.is_none());
}
#[test]
fn asm_await_mem_plans_returns_canned_plans_after_thread_join() {
let canned = Vec::<Plan>::new();
let expected_len = canned.len();
let mo_handle = std::thread::spawn(move || Ok(AsmRunnerMO::new(canned)));
let mut backend = BackendArtifacts::Asm { mo: Some(mo_handle), rh: None };
let (plans, gpu_mops_used_bytes) =
backend.await_mem_plans().expect("await_mem_plans on Asm");
assert_eq!(plans.len(), expected_len);
assert!(gpu_mops_used_bytes.is_none());
}
#[test]
fn asm_await_mem_plans_errs_on_double_take() {
let mo_handle = std::thread::spawn(move || Ok(AsmRunnerMO::new(Vec::new())));
let mut backend = BackendArtifacts::Asm { mo: Some(mo_handle), rh: None };
backend.await_mem_plans().expect("first call OK");
let err = backend.await_mem_plans().expect_err("second call must err");
assert!(err.to_string().contains("already consumed"));
}
#[test]
fn asm_await_mem_plans_propagates_runner_error() {
let mo_handle = std::thread::spawn(|| -> ExecutorResult<AsmRunnerMO> {
Err(ExecutorError::AsmBackend("boom".to_string()))
});
let mut backend = BackendArtifacts::Asm { mo: Some(mo_handle), rh: None };
let err = backend.await_mem_plans().expect_err("runner Err must propagate");
assert!(err.to_string().contains("MO runner failed"));
assert!(err.to_string().contains("boom"));
}
#[test]
fn asm_await_rom_histogram_returns_some_after_join() {
let rh_handle = std::thread::spawn(|| Ok(AsmRunnerRH::new(AsmRHData::new(0, Vec::new()))));
let mut backend = BackendArtifacts::Asm { mo: None, rh: Some(rh_handle) };
let rh = backend.await_rom_histogram().expect("await_rom_histogram on Asm");
assert!(rh.is_some());
}
#[test]
fn asm_await_rom_histogram_none_when_not_present() {
let mut backend = BackendArtifacts::Asm { mo: None, rh: None };
let rh = backend.await_rom_histogram().expect("await_rom_histogram with rh=None");
assert!(rh.is_none());
}
#[test]
fn asm_await_rom_histogram_double_take_yields_none() {
let rh_handle = std::thread::spawn(|| Ok(AsmRunnerRH::new(AsmRHData::new(0, Vec::new()))));
let mut backend = BackendArtifacts::Asm { mo: None, rh: Some(rh_handle) };
backend.await_rom_histogram().expect("first call OK");
let second = backend.await_rom_histogram().expect("second call OK (None)");
assert!(second.is_none());
}
}