use std::collections::HashMap;
use proofman_fields::PrimeField64;
use zisk_common::{
plan, BusDeviceMetrics, ChunkId, CollectSkipper, InstCount, InstanceType, Metrics, Plan,
Planner, SegmentId,
};
use zisk_pil::{JumpDestTrace, JUMP_DEST_AIR_IDS, ZISK_AIRGROUP_ID};
use crate::{JumpDestCheckPoint, JumpDestCounterInputGen};
#[derive(Default)]
pub struct JumpDestPlanner<F> {
_marker: std::marker::PhantomData<F>,
}
impl<F: PrimeField64> JumpDestPlanner<F> {
pub fn new() -> Self {
Self::default()
}
}
impl<F: PrimeField64> Planner for JumpDestPlanner<F> {
fn plan(&self, counters: Vec<(ChunkId, Box<dyn BusDeviceMetrics>)>) -> Vec<Plan> {
let counts: Vec<InstCount> = counters
.iter()
.map(|(chunk_id, counter)| {
let rows = Metrics::as_any(&**counter)
.downcast_ref::<JumpDestCounterInputGen>()
.expect("JumpDestPlanner got a counter that is not a JumpDest one")
.rows;
InstCount::new(*chunk_id, rows as u64)
})
.collect();
let segments = plan(&counts, JumpDestTrace::<usize>::NUM_ROWS as u64);
let last = segments.len().saturating_sub(1);
segments
.into_iter()
.enumerate()
.map(|(segment, (check_point, collect_info))| {
let chunks: HashMap<ChunkId, (u64, CollectSkipper)> = collect_info;
let meta = JumpDestCheckPoint {
last_chunk: chunks.keys().max().copied(),
is_last_segment: segment == last,
chunks,
};
Plan::new(
ZISK_AIRGROUP_ID,
JUMP_DEST_AIR_IDS[0],
Some(SegmentId(segment)),
InstanceType::Instance,
check_point,
Some(Box::new(meta)),
)
})
.collect()
}
}