openvm_sha2_circuit/sha2_chips/
trace.rs

1use std::{
2    borrow::BorrowMut,
3    mem::transmute,
4    slice::{from_raw_parts, from_raw_parts_mut},
5};
6
7use openvm_circuit::{
8    arch::{
9        CustomBorrow, ExecutionError, MultiRowLayout, MultiRowMetadata, PreflightExecutor,
10        RecordArena, SizedRecord, VmStateMut,
11    },
12    system::memory::{
13        offline_checker::{MemoryReadAuxRecord, MemoryWriteBytesAuxRecord},
14        online::TracingMemory,
15    },
16};
17use openvm_circuit_primitives::AlignedBytesBorrow;
18use openvm_instructions::{
19    instruction::Instruction,
20    program::DEFAULT_PC_STEP,
21    riscv::{RV32_MEMORY_AS, RV32_REGISTER_AS},
22    LocalOpcode,
23};
24use openvm_rv32im_circuit::adapters::{tracing_read, tracing_write};
25use openvm_sha2_air::{Sha256Config, Sha2Variant, Sha384Config, Sha512Config};
26use openvm_stark_backend::p3_field::PrimeField32;
27
28use crate::{
29    Sha2Config, Sha2MainChipConfig, Sha2VmExecutor, SHA2_READ_SIZE, SHA2_REGISTER_READS,
30    SHA2_WRITE_SIZE,
31};
32
33#[derive(Clone, Copy)]
34pub struct Sha2Metadata {
35    pub variant: Sha2Variant,
36}
37
38impl MultiRowMetadata for Sha2Metadata {
39    #[inline(always)]
40    fn get_num_rows(&self) -> usize {
41        // The size of the record arena will be height * Sha2MainAir::width() * num_rows.
42        // We will not use the record arena's buffer for either chip's trace, so we just
43        // need to ensure that the record arena is large enough to store all the records.
44        // The size of Sha2RecordMut (in bytes) is less than Sha2MainAir::width() * size_of::<F>(),
45        // for all SHA-2 variants. Therefore, we can set num_rows = 1.
46        1
47    }
48}
49
50pub(crate) type Sha2RecordLayout = MultiRowLayout<Sha2Metadata>;
51
52#[repr(C)]
53#[derive(AlignedBytesBorrow, Debug, Clone)]
54pub struct Sha2RecordHeader {
55    pub variant: Sha2Variant,
56    pub from_pc: u32,
57    pub timestamp: u32,
58    pub dst_reg_ptr: u32,
59    pub state_reg_ptr: u32,
60    pub input_reg_ptr: u32,
61    pub dst_ptr: u32,
62    pub state_ptr: u32,
63    pub input_ptr: u32,
64
65    pub register_reads_aux: [MemoryReadAuxRecord; SHA2_REGISTER_READS],
66}
67
68pub struct Sha2RecordMut<'a> {
69    pub inner: &'a mut Sha2RecordHeader,
70
71    pub message_bytes: &'a mut [u8],
72    pub prev_state: &'a mut [u8], // little-endian words
73    pub new_state: &'a mut [u8],  // little-endian words
74
75    pub input_reads_aux: &'a mut [MemoryReadAuxRecord],
76    pub state_reads_aux: &'a mut [MemoryReadAuxRecord],
77    pub write_aux: &'a mut [MemoryWriteBytesAuxRecord<SHA2_WRITE_SIZE>],
78}
79
80impl<'a> CustomBorrow<'a, Sha2RecordMut<'a>, Sha2RecordLayout> for [u8] {
81    fn custom_borrow(&'a mut self, layout: Sha2RecordLayout) -> Sha2RecordMut<'a> {
82        // SAFETY:
83        // - Caller guarantees through the layout that self has sufficient length for all splits and
84        //   constants are guaranteed <= self.len() by layout precondition
85
86        let (header_slice, rest) =
87            unsafe { self.split_at_mut_unchecked(size_of::<Sha2RecordHeader>()) };
88        let record_header: &mut Sha2RecordHeader = header_slice.borrow_mut();
89
90        let dims = Sha2PreComputeDims::new(layout.metadata.variant);
91
92        let (message_bytes, rest) = unsafe { rest.split_at_mut_unchecked(dims.input_size) };
93        let (prev_state, rest) = unsafe { rest.split_at_mut_unchecked(dims.state_size) };
94        let (new_state, rest) = unsafe { rest.split_at_mut_unchecked(dims.state_size) };
95
96        let (input_reads_aux, rest) = unsafe { align_to_mut_at(rest, dims.input_reads) };
97        let (state_reads_aux, rest) = unsafe { align_to_mut_at(rest, dims.state_reads) };
98        let (write_aux, _) = unsafe { align_to_mut_at(rest, dims.state_writes) };
99
100        Sha2RecordMut {
101            inner: record_header,
102            message_bytes,
103            prev_state,
104            new_state,
105            input_reads_aux,
106            state_reads_aux,
107            write_aux,
108        }
109    }
110
111    unsafe fn extract_layout(&self) -> Sha2RecordLayout {
112        let (variant, _) = unsafe { align_to_at(self, 1) };
113        let variant = variant[0];
114        Sha2RecordLayout {
115            metadata: Sha2Metadata { variant },
116        }
117    }
118}
119
120unsafe fn align_to_mut_at<T>(slice: &mut [u8], offset: usize) -> (&mut [T], &mut [u8]) {
121    let (_, items, rest) = unsafe { slice.align_to_mut::<T>() };
122    let (items, items_rest) = unsafe { items.split_at_mut_unchecked(offset) };
123    let rest = unsafe {
124        let items_rest: &mut [u8] = transmute(items_rest);
125        from_raw_parts_mut(
126            items_rest.as_mut_ptr(),
127            items_rest.len() * size_of::<T>() + rest.len(),
128        )
129    };
130    (items, rest)
131}
132
133unsafe fn align_to_at<T>(slice: &[u8], offset: usize) -> (&[T], &[u8]) {
134    let (_, items, rest) = unsafe { slice.align_to::<T>() };
135    let (items, items_rest) = unsafe { items.split_at_unchecked(offset) };
136    let rest = unsafe {
137        let items_rest: &[u8] = transmute(items_rest);
138        from_raw_parts(
139            items_rest.as_ptr(),
140            items_rest.len() * size_of::<T>() + rest.len(),
141        )
142    };
143    (items, rest)
144}
145
146impl SizedRecord<Sha2RecordLayout> for Sha2RecordMut<'_> {
147    fn size(layout: &Sha2RecordLayout) -> usize {
148        let header_size = size_of::<Sha2RecordHeader>();
149        let dims = Sha2PreComputeDims::new(layout.metadata.variant);
150        let mut total_len = header_size
151            + dims.input_size  // input
152            + dims.state_size  // prev_state
153            + dims.state_size; // new_state
154
155        total_len = total_len.next_multiple_of(align_of::<MemoryReadAuxRecord>());
156        total_len += dims.input_reads * size_of::<MemoryReadAuxRecord>();
157
158        total_len = total_len.next_multiple_of(align_of::<MemoryReadAuxRecord>());
159        total_len += dims.state_reads * size_of::<MemoryReadAuxRecord>();
160
161        total_len =
162            total_len.next_multiple_of(align_of::<MemoryWriteBytesAuxRecord<SHA2_WRITE_SIZE>>());
163        total_len += dims.state_writes * size_of::<MemoryWriteBytesAuxRecord<SHA2_WRITE_SIZE>>();
164
165        total_len
166    }
167
168    fn alignment(_layout: &Sha2RecordLayout) -> usize {
169        align_of::<Sha2RecordHeader>() // 4-byte alignment
170    }
171}
172
173// This is needed in CustomBorrow trait to convert the Sha2Variant that we read from the buffer
174// into appropriate dimensions for the record.
175struct Sha2PreComputeDims {
176    state_size: usize,
177    input_size: usize,
178    input_reads: usize,
179    state_reads: usize,
180    state_writes: usize,
181}
182
183impl Sha2PreComputeDims {
184    fn new(variant: Sha2Variant) -> Self {
185        match variant {
186            Sha2Variant::Sha256 => Self {
187                state_size: Sha256Config::STATE_BYTES,
188                input_size: Sha256Config::BLOCK_BYTES,
189                input_reads: Sha256Config::BLOCK_READS,
190                state_reads: Sha256Config::STATE_READS,
191                state_writes: Sha256Config::STATE_WRITES,
192            },
193            Sha2Variant::Sha512 => Self {
194                state_size: Sha512Config::STATE_BYTES,
195                input_size: Sha512Config::BLOCK_BYTES,
196                input_reads: Sha512Config::BLOCK_READS,
197                state_reads: Sha512Config::STATE_READS,
198                state_writes: Sha512Config::STATE_WRITES,
199            },
200            Sha2Variant::Sha384 => Self {
201                state_size: Sha384Config::STATE_BYTES,
202                input_size: Sha384Config::BLOCK_BYTES,
203                input_reads: Sha384Config::BLOCK_READS,
204                state_reads: Sha384Config::STATE_READS,
205                state_writes: Sha384Config::STATE_WRITES,
206            },
207        }
208    }
209}
210
211impl<F, RA, C: Sha2Config> PreflightExecutor<F, RA> for Sha2VmExecutor<C>
212where
213    F: PrimeField32,
214    // for<'buf> RA: RecordArena<'buf, Sha2RecordLayout, Sha2RecordMut<'buf>>,
215    for<'buf> RA: RecordArena<'buf, Sha2RecordLayout, Sha2RecordMut<'buf>>,
216{
217    fn get_opcode_name(&self, _: usize) -> String {
218        format!("{:?}", C::OPCODE)
219    }
220
221    fn execute(
222        &self,
223        state: VmStateMut<F, TracingMemory, RA>,
224        instruction: &Instruction<F>,
225    ) -> Result<(), ExecutionError> {
226        let &Instruction {
227            opcode,
228            a,
229            b,
230            c,
231            d,
232            e,
233            ..
234        } = instruction;
235        debug_assert_eq!(opcode, C::OPCODE.global_opcode());
236        debug_assert_eq!(d.as_canonical_u32(), RV32_REGISTER_AS);
237        debug_assert_eq!(e.as_canonical_u32(), RV32_MEMORY_AS);
238
239        let record = state.ctx.alloc(Sha2RecordLayout::new(Sha2Metadata {
240            variant: C::VARIANT,
241        }));
242
243        record.inner.variant = C::VARIANT;
244        record.inner.from_pc = *state.pc;
245        record.inner.timestamp = state.memory.timestamp();
246        record.inner.dst_reg_ptr = a.as_canonical_u32();
247        record.inner.state_reg_ptr = b.as_canonical_u32();
248        record.inner.input_reg_ptr = c.as_canonical_u32();
249
250        record.inner.dst_ptr = u32::from_le_bytes(tracing_read::<SHA2_READ_SIZE>(
251            state.memory,
252            RV32_REGISTER_AS,
253            record.inner.dst_reg_ptr,
254            &mut record.inner.register_reads_aux[0].prev_timestamp,
255        ));
256        record.inner.state_ptr = u32::from_le_bytes(tracing_read::<SHA2_READ_SIZE>(
257            state.memory,
258            RV32_REGISTER_AS,
259            record.inner.state_reg_ptr,
260            &mut record.inner.register_reads_aux[1].prev_timestamp,
261        ));
262        record.inner.input_ptr = u32::from_le_bytes(tracing_read::<SHA2_READ_SIZE>(
263            state.memory,
264            RV32_REGISTER_AS,
265            record.inner.input_reg_ptr,
266            &mut record.inner.register_reads_aux[2].prev_timestamp,
267        ));
268
269        debug_assert!(
270            record.inner.dst_ptr as usize + C::STATE_BYTES <= (1 << self.pointer_max_bits)
271        );
272        debug_assert!(
273            record.inner.state_ptr as usize + C::STATE_BYTES <= (1 << self.pointer_max_bits)
274        );
275        debug_assert!(
276            record.inner.input_ptr as usize + C::BLOCK_BYTES <= (1 << self.pointer_max_bits)
277        );
278
279        for idx in 0..C::BLOCK_READS {
280            let read = tracing_read::<SHA2_READ_SIZE>(
281                state.memory,
282                RV32_MEMORY_AS,
283                record.inner.input_ptr + (idx * SHA2_READ_SIZE) as u32,
284                &mut record.input_reads_aux[idx].prev_timestamp,
285            );
286            record.message_bytes[idx * SHA2_READ_SIZE..(idx + 1) * SHA2_READ_SIZE]
287                .copy_from_slice(&read);
288        }
289
290        for idx in 0..C::STATE_READS {
291            let read = tracing_read::<SHA2_READ_SIZE>(
292                state.memory,
293                RV32_MEMORY_AS,
294                record.inner.state_ptr + (idx * SHA2_READ_SIZE) as u32,
295                &mut record.state_reads_aux[idx].prev_timestamp,
296            );
297            record.prev_state[idx * SHA2_READ_SIZE..(idx + 1) * SHA2_READ_SIZE]
298                .copy_from_slice(&read);
299        }
300
301        record.new_state.copy_from_slice(record.prev_state);
302        C::compress(record.new_state, record.message_bytes);
303
304        for idx in 0..C::STATE_WRITES {
305            tracing_write::<SHA2_WRITE_SIZE>(
306                state.memory,
307                RV32_MEMORY_AS,
308                record.inner.dst_ptr + (idx * SHA2_WRITE_SIZE) as u32,
309                record.new_state[idx * SHA2_WRITE_SIZE..(idx + 1) * SHA2_WRITE_SIZE]
310                    .try_into()
311                    .unwrap(),
312                &mut record.write_aux[idx].prev_timestamp,
313                &mut record.write_aux[idx].prev_data,
314            );
315        }
316
317        *state.pc = state.pc.wrapping_add(DEFAULT_PC_STEP);
318        Ok(())
319    }
320}