use std::{collections::BTreeMap, fmt::Debug, marker::PhantomData, sync::Arc};
use getset::{Getters, MutGetters};
use openvm_circuit_primitives::{
assert_less_than::{AssertLtSubAir, LessThanAuxCols},
var_range::{
SharedVariableRangeCheckerChip, VariableRangeCheckerBus, VariableRangeCheckerChip,
},
Chip, TraceSubRowGenerator,
};
use openvm_cpu_backend::CpuBackend;
use openvm_stark_backend::{
interaction::PermutationCheckBus, p3_field::PrimeField32, p3_util::log2_strict_usize,
prover::AirProvingContext, StarkProtocolConfig,
};
use serde::{Deserialize, Serialize};
use self::interface::MemoryInterface;
use super::AddressMap;
use crate::{
arch::{MemoryConfig, VmField, DEFAULT_BLOCK_SIZE},
system::{
memory::{
dimensions::MemoryDimensions,
merkle::MemoryMerkleChip,
offline_checker::{MemoryBaseAuxCols, MemoryBridge, MemoryBus, AUX_LEN},
persistent::{group_touched_memory_by_chunk, PersistentBoundaryChip},
},
poseidon2::Poseidon2PeripheryChip,
TouchedMemory,
},
};
pub mod dimensions;
pub mod interface;
pub const CHUNK: usize = 8;
pub const MERKLE_AIR_OFFSET: usize = 1;
pub const BOUNDARY_AIR_OFFSET: usize = 0;
pub type MemoryImage = AddressMap;
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct TimestampedValues<T, const N: usize> {
pub timestamp: u32,
pub values: [T; N],
}
pub type TimestampedEquipartition<F, const N: usize> = Vec<((u32, u32), TimestampedValues<F, N>)>;
pub type Equipartition<F, const N: usize> = BTreeMap<(u32, u32), [F; N]>;
#[derive(Getters, MutGetters)]
pub struct MemoryController<F: VmField> {
pub memory_bus: MemoryBus,
pub interface_chip: MemoryInterface<F>,
pub range_checker: SharedVariableRangeCheckerChip,
pub(crate) memory_config: MemoryConfig,
range_checker_bus: VariableRangeCheckerBus,
pub(crate) hasher_chip: Option<Arc<Poseidon2PeripheryChip<F>>>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct PersistentMemoryTraceHeights {
boundary: usize,
merkle: usize,
}
impl PersistentMemoryTraceHeights {
pub fn from_slice(heights: &[u32]) -> Self {
Self {
boundary: heights[0] as usize,
merkle: heights[1] as usize,
}
}
}
impl<F: VmField> MemoryController<F> {
pub fn with_persistent_memory(
memory_bus: MemoryBus,
mem_config: MemoryConfig,
range_checker: SharedVariableRangeCheckerChip,
merkle_bus: PermutationCheckBus,
compression_bus: PermutationCheckBus,
hasher_chip: Arc<Poseidon2PeripheryChip<F>>,
) -> Self {
let memory_dims = MemoryDimensions {
addr_space_height: mem_config.addr_space_height,
address_height: mem_config.pointer_max_bits - log2_strict_usize(CHUNK),
};
let range_checker_bus = range_checker.bus();
let interface_chip = MemoryInterface {
boundary_chip: PersistentBoundaryChip::new(memory_bus, merkle_bus, compression_bus),
merkle_chip: MemoryMerkleChip::new(memory_dims, merkle_bus, compression_bus),
initial_memory: AddressMap::from_mem_config(&mem_config),
};
Self {
memory_bus,
interface_chip,
memory_config: mem_config,
range_checker,
range_checker_bus,
hasher_chip: Some(hasher_chip),
}
}
pub fn memory_config(&self) -> &MemoryConfig {
&self.memory_config
}
pub(crate) fn set_override_trace_heights(&mut self, overridden_heights: &[u32]) {
let oh = PersistentMemoryTraceHeights::from_slice(overridden_heights);
self.interface_chip
.boundary_chip
.set_overridden_height(oh.boundary);
self.interface_chip
.merkle_chip
.set_overridden_height(oh.merkle);
}
pub(crate) fn set_initial_memory(&mut self, memory: AddressMap) {
self.interface_chip.initial_memory = memory;
}
pub fn memory_bridge(&self) -> MemoryBridge {
MemoryBridge::new(
self.memory_bus,
self.memory_config().timestamp_max_bits,
self.range_checker_bus,
)
}
pub fn helper(&self) -> SharedMemoryHelper<F> {
let range_bus = self.range_checker.bus();
SharedMemoryHelper {
range_checker: self.range_checker.clone(),
timestamp_lt_air: AssertLtSubAir::new(
range_bus,
self.memory_config().timestamp_max_bits,
),
_marker: Default::default(),
}
}
pub fn generate_proving_ctx<SC: StarkProtocolConfig<F = F>>(
&mut self,
touched_memory: TouchedMemory<F>,
) -> Vec<AirProvingContext<CpuBackend<SC>>> {
let final_memory = touched_memory;
let MemoryInterface {
boundary_chip,
merkle_chip,
initial_memory,
} = &mut self.interface_chip;
let hasher = self.hasher_chip.as_ref().unwrap();
boundary_chip.finalize(initial_memory, &final_memory, hasher.as_ref());
let final_memory_values: Equipartition<F, CHUNK> =
group_touched_memory_by_chunk(&final_memory)
.into_iter()
.map(|((addr_space, chunk_label), blocks)| {
let chunk_ptr = chunk_label * CHUNK as u32;
let mut values = std::array::from_fn(|i| unsafe {
initial_memory.get_f::<F>(addr_space, chunk_ptr + i as u32)
});
for (block_idx, _, block_values) in blocks {
for (i, val) in block_values.into_iter().enumerate() {
values[block_idx * DEFAULT_BLOCK_SIZE + i] = val;
}
}
((addr_space, chunk_ptr), values)
})
.collect();
merkle_chip.finalize(initial_memory, &final_memory_values, hasher.as_ref());
vec![
boundary_chip.generate_proving_ctx(()),
merkle_chip.generate_proving_ctx(),
]
}
pub fn num_airs(&self) -> usize {
2
}
}
#[derive(Clone)]
pub struct SharedMemoryHelper<F> {
pub(crate) range_checker: SharedVariableRangeCheckerChip,
pub(crate) timestamp_lt_air: AssertLtSubAir,
pub(crate) _marker: PhantomData<F>,
}
impl<F> SharedMemoryHelper<F> {
pub fn new(range_checker: SharedVariableRangeCheckerChip, timestamp_max_bits: usize) -> Self {
let timestamp_lt_air = AssertLtSubAir::new(range_checker.bus(), timestamp_max_bits);
Self {
range_checker,
timestamp_lt_air,
_marker: PhantomData,
}
}
}
pub struct MemoryAuxColsFactory<'a, F> {
pub(crate) range_checker: &'a VariableRangeCheckerChip,
pub(crate) timestamp_lt_air: AssertLtSubAir,
pub(crate) _marker: PhantomData<F>,
}
impl<F: PrimeField32> MemoryAuxColsFactory<'_, F> {
pub fn fill(&self, prev_timestamp: u32, timestamp: u32, buffer: &mut MemoryBaseAuxCols<F>) {
self.generate_timestamp_lt(prev_timestamp, timestamp, &mut buffer.timestamp_lt_aux);
buffer.prev_timestamp = F::from_u32(prev_timestamp);
}
pub fn fill_zero(&self, buffer: &mut MemoryBaseAuxCols<F>) {
*buffer = unsafe { std::mem::zeroed() };
}
fn generate_timestamp_lt(
&self,
prev_timestamp: u32,
timestamp: u32,
buffer: &mut LessThanAuxCols<F, AUX_LEN>,
) {
debug_assert!(
prev_timestamp < timestamp,
"prev_timestamp {prev_timestamp} >= timestamp {timestamp}"
);
self.timestamp_lt_air.generate_subrow(
(self.range_checker, prev_timestamp, timestamp),
&mut buffer.lower_decomp,
);
}
}
impl<F> SharedMemoryHelper<F> {
pub fn as_borrowed(&self) -> MemoryAuxColsFactory<'_, F> {
MemoryAuxColsFactory {
range_checker: self.range_checker.as_ref(),
timestamp_lt_air: self.timestamp_lt_air,
_marker: PhantomData,
}
}
}