openvm_circuit/system/memory/controller/
mod.rs1use std::{collections::BTreeMap, fmt::Debug, marker::PhantomData, sync::Arc};
3
4use getset::{Getters, MutGetters};
5use openvm_circuit_primitives::{
6 assert_less_than::{AssertLtSubAir, LessThanAuxCols},
7 var_range::{
8 SharedVariableRangeCheckerChip, VariableRangeCheckerBus, VariableRangeCheckerChip,
9 },
10 Chip, TraceSubRowGenerator,
11};
12use openvm_cpu_backend::CpuBackend;
13use openvm_stark_backend::{
14 interaction::PermutationCheckBus, p3_field::PrimeField32, p3_util::log2_strict_usize,
15 prover::AirProvingContext, StarkProtocolConfig,
16};
17use serde::{Deserialize, Serialize};
18
19use self::interface::MemoryInterface;
20use super::AddressMap;
21use crate::{
22 arch::{MemoryConfig, VmField, DEFAULT_BLOCK_SIZE},
23 system::{
24 memory::{
25 dimensions::MemoryDimensions,
26 merkle::MemoryMerkleChip,
27 offline_checker::{MemoryBaseAuxCols, MemoryBridge, MemoryBus, AUX_LEN},
28 persistent::{group_touched_memory_by_chunk, PersistentBoundaryChip},
29 },
30 poseidon2::Poseidon2PeripheryChip,
31 TouchedMemory,
32 },
33};
34
35pub mod dimensions;
36pub mod interface;
37
38pub const CHUNK: usize = 8;
39
40pub const MERKLE_AIR_OFFSET: usize = 1;
42pub const BOUNDARY_AIR_OFFSET: usize = 0;
44
45pub type MemoryImage = AddressMap;
46
47#[repr(C)]
48#[derive(Clone, Copy, Debug, PartialEq, Eq)]
49pub struct TimestampedValues<T, const N: usize> {
50 pub timestamp: u32,
51 pub values: [T; N],
52}
53
54pub type TimestampedEquipartition<F, const N: usize> = Vec<((u32, u32), TimestampedValues<F, N>)>;
59
60pub type Equipartition<F, const N: usize> = BTreeMap<(u32, u32), [F; N]>;
67
68#[derive(Getters, MutGetters)]
69pub struct MemoryController<F: VmField> {
70 pub memory_bus: MemoryBus,
71 pub interface_chip: MemoryInterface<F>,
72 pub range_checker: SharedVariableRangeCheckerChip,
73 pub(crate) memory_config: MemoryConfig,
74 range_checker_bus: VariableRangeCheckerBus,
76 pub(crate) hasher_chip: Option<Arc<Poseidon2PeripheryChip<F>>>,
77}
78
79#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
80pub struct PersistentMemoryTraceHeights {
81 boundary: usize,
82 merkle: usize,
83}
84impl PersistentMemoryTraceHeights {
85 pub fn from_slice(heights: &[u32]) -> Self {
87 Self {
88 boundary: heights[0] as usize,
89 merkle: heights[1] as usize,
90 }
91 }
92}
93
94impl<F: VmField> MemoryController<F> {
95 pub fn with_persistent_memory(
99 memory_bus: MemoryBus,
100 mem_config: MemoryConfig,
101 range_checker: SharedVariableRangeCheckerChip,
102 merkle_bus: PermutationCheckBus,
103 compression_bus: PermutationCheckBus,
104 hasher_chip: Arc<Poseidon2PeripheryChip<F>>,
105 ) -> Self {
106 let memory_dims = MemoryDimensions {
107 addr_space_height: mem_config.addr_space_height,
108 address_height: mem_config.pointer_max_bits - log2_strict_usize(CHUNK),
109 };
110 let range_checker_bus = range_checker.bus();
111 let interface_chip = MemoryInterface {
112 boundary_chip: PersistentBoundaryChip::new(memory_bus, merkle_bus, compression_bus),
113 merkle_chip: MemoryMerkleChip::new(memory_dims, merkle_bus, compression_bus),
114 initial_memory: AddressMap::from_mem_config(&mem_config),
115 };
116 Self {
117 memory_bus,
118 interface_chip,
119 memory_config: mem_config,
120 range_checker,
121 range_checker_bus,
122 hasher_chip: Some(hasher_chip),
123 }
124 }
125
126 pub fn memory_config(&self) -> &MemoryConfig {
127 &self.memory_config
128 }
129
130 pub(crate) fn set_override_trace_heights(&mut self, overridden_heights: &[u32]) {
131 let oh = PersistentMemoryTraceHeights::from_slice(overridden_heights);
132 self.interface_chip
133 .boundary_chip
134 .set_overridden_height(oh.boundary);
135 self.interface_chip
136 .merkle_chip
137 .set_overridden_height(oh.merkle);
138 }
139
140 pub(crate) fn set_initial_memory(&mut self, memory: AddressMap) {
143 self.interface_chip.initial_memory = memory;
144 }
145
146 pub fn memory_bridge(&self) -> MemoryBridge {
147 MemoryBridge::new(
148 self.memory_bus,
149 self.memory_config().timestamp_max_bits,
150 self.range_checker_bus,
151 )
152 }
153
154 pub fn helper(&self) -> SharedMemoryHelper<F> {
155 let range_bus = self.range_checker.bus();
156 SharedMemoryHelper {
157 range_checker: self.range_checker.clone(),
158 timestamp_lt_air: AssertLtSubAir::new(
159 range_bus,
160 self.memory_config().timestamp_max_bits,
161 ),
162 _marker: Default::default(),
163 }
164 }
165
166 pub fn generate_proving_ctx<SC: StarkProtocolConfig<F = F>>(
171 &mut self,
172 touched_memory: TouchedMemory<F>,
173 ) -> Vec<AirProvingContext<CpuBackend<SC>>> {
174 let final_memory = touched_memory;
175 let MemoryInterface {
176 boundary_chip,
177 merkle_chip,
178 initial_memory,
179 } = &mut self.interface_chip;
180
181 let hasher = self.hasher_chip.as_ref().unwrap();
182 boundary_chip.finalize(initial_memory, &final_memory, hasher.as_ref());
183
184 let final_memory_values: Equipartition<F, CHUNK> =
187 group_touched_memory_by_chunk(&final_memory)
188 .into_iter()
189 .map(|((addr_space, chunk_label), blocks)| {
190 let chunk_ptr = chunk_label * CHUNK as u32;
191 let mut values = std::array::from_fn(|i| unsafe {
192 initial_memory.get_f::<F>(addr_space, chunk_ptr + i as u32)
193 });
194 for (block_idx, _, block_values) in blocks {
195 for (i, val) in block_values.into_iter().enumerate() {
196 values[block_idx * DEFAULT_BLOCK_SIZE + i] = val;
197 }
198 }
199 ((addr_space, chunk_ptr), values)
200 })
201 .collect();
202 merkle_chip.finalize(initial_memory, &final_memory_values, hasher.as_ref());
203
204 vec![
205 boundary_chip.generate_proving_ctx(()),
206 merkle_chip.generate_proving_ctx(),
207 ]
208 }
209
210 pub fn num_airs(&self) -> usize {
212 2
213 }
214}
215
216#[derive(Clone)]
218pub struct SharedMemoryHelper<F> {
219 pub(crate) range_checker: SharedVariableRangeCheckerChip,
220 pub(crate) timestamp_lt_air: AssertLtSubAir,
221 pub(crate) _marker: PhantomData<F>,
222}
223
224impl<F> SharedMemoryHelper<F> {
225 pub fn new(range_checker: SharedVariableRangeCheckerChip, timestamp_max_bits: usize) -> Self {
226 let timestamp_lt_air = AssertLtSubAir::new(range_checker.bus(), timestamp_max_bits);
227 Self {
228 range_checker,
229 timestamp_lt_air,
230 _marker: PhantomData,
231 }
232 }
233}
234
235pub struct MemoryAuxColsFactory<'a, F> {
238 pub(crate) range_checker: &'a VariableRangeCheckerChip,
239 pub(crate) timestamp_lt_air: AssertLtSubAir,
240 pub(crate) _marker: PhantomData<F>,
241}
242
243impl<F: PrimeField32> MemoryAuxColsFactory<'_, F> {
244 pub fn fill(&self, prev_timestamp: u32, timestamp: u32, buffer: &mut MemoryBaseAuxCols<F>) {
246 self.generate_timestamp_lt(prev_timestamp, timestamp, &mut buffer.timestamp_lt_aux);
247 buffer.prev_timestamp = F::from_u32(prev_timestamp);
250 }
251
252 pub fn fill_zero(&self, buffer: &mut MemoryBaseAuxCols<F>) {
255 *buffer = unsafe { std::mem::zeroed() };
256 }
257
258 fn generate_timestamp_lt(
259 &self,
260 prev_timestamp: u32,
261 timestamp: u32,
262 buffer: &mut LessThanAuxCols<F, AUX_LEN>,
263 ) {
264 debug_assert!(
265 prev_timestamp < timestamp,
266 "prev_timestamp {prev_timestamp} >= timestamp {timestamp}"
267 );
268 self.timestamp_lt_air.generate_subrow(
269 (self.range_checker, prev_timestamp, timestamp),
270 &mut buffer.lower_decomp,
271 );
272 }
273}
274
275impl<F> SharedMemoryHelper<F> {
276 pub fn as_borrowed(&self) -> MemoryAuxColsFactory<'_, F> {
277 MemoryAuxColsFactory {
278 range_checker: self.range_checker.as_ref(),
279 timestamp_lt_air: self.timestamp_lt_air,
280 _marker: PhantomData,
281 }
282 }
283}