openvm_circuit/system/memory/controller/
mod.rs

1//! [MemoryController] can be considered as the Memory Chip Complex for the CPU Backend.
2use 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
40/// The offset of the Merkle AIR in AIRs of MemoryController.
41pub const MERKLE_AIR_OFFSET: usize = 1;
42/// The offset of the boundary AIR in AIRs of MemoryController.
43pub 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
54/// A sorted equipartition of memory, with timestamps and values.
55///
56/// The "key" is a pair `(address_space, label)`, where `label` is the index of the block in the
57/// partition. I.e., the starting address of the block is `(address_space, label * N)`.
58pub type TimestampedEquipartition<F, const N: usize> = Vec<((u32, u32), TimestampedValues<F, N>)>;
59
60/// An equipartition of memory values.
61///
62/// The key is a pair `(address_space, label)`, where `label` is the index of the block in the
63/// partition. I.e., the starting address of the block is `(address_space, label * N)`.
64///
65/// If a key is not present in the map, then the block is uninitialized (and therefore zero).
66pub 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    // Store separately to avoid smart pointer reference each time
75    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    /// `heights` must consist of only memory trace heights, in order of AIR IDs.
86    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    /// Creates a new memory controller for persistent memory.
96    ///
97    /// Call `set_initial_memory` to set the initial memory state after construction.
98    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    /// This only sets the initial memory image for the boundary and merkle tree chips.
141    /// Tracing memory should be set separately.
142    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    // @dev: Memory is complicated and allowed to break all the rules (e.g., 1 arena per chip) and
167    // there's no need for any memory chip to implement the Chip trait. We do it when convenient,
168    // but all that matters is that you can tracegen all the trace matrices for the memory AIRs
169    // _somehow_.
170    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        // Rechunk DEFAULT_BLOCK_SIZE blocks into CHUNK-sized blocks for merkle_chip
185        // Note: Equipartition key is (addr_space, ptr) where ptr is the starting pointer
186        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    /// Return the number of AIRs in the memory controller.
211    pub fn num_airs(&self) -> usize {
212        2
213    }
214}
215
216/// Owned version of [MemoryAuxColsFactory].
217#[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
235/// A helper for generating trace values in auxiliary memory columns related to the offline memory
236/// argument.
237pub 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    /// Fill the trace assuming `prev_timestamp` is already provided in `buffer`.
245    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        // Safety: even if prev_timestamp were obtained by transmute_ref from
248        // `buffer.prev_timestamp`, this should still work because it is a direct assignment
249        buffer.prev_timestamp = F::from_u32(prev_timestamp);
250    }
251
252    /// # Safety
253    /// We assume that `F::ZERO` has underlying memory equivalent to `mem::zeroed()`.
254    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}