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 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], pub new_state: &'a mut [u8], 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 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 + dims.state_size + dims.state_size; 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>() }
171}
172
173struct 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>>,
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}