openvm_deferral_circuit/output/
execution.rs1use std::{
2 array::from_fn,
3 borrow::{Borrow, BorrowMut},
4 slice::from_raw_parts,
5};
6
7use openvm_circuit::{arch::*, system::memory::online::GuestMemory};
8use openvm_circuit_primitives::AlignedBytesBorrow;
9use openvm_deferral_transpiler::DeferralOpcode;
10use openvm_instructions::{
11 instruction::Instruction,
12 program::DEFAULT_PC_STEP,
13 riscv::{RV32_MEMORY_AS, RV32_REGISTER_AS},
14 LocalOpcode,
15};
16use openvm_stark_backend::p3_field::PrimeField32;
17use openvm_stark_sdk::config::baby_bear_poseidon2::DIGEST_SIZE;
18
19use super::DeferralOutputExecutor;
20use crate::{
21 utils::{
22 join_memory_ops, memory_op_chunk, split_output, DIGEST_MEMORY_OPS, OUTPUT_TOTAL_BYTES,
23 OUTPUT_TOTAL_MEMORY_OPS,
24 },
25 OUTPUT_AIR_REL_IDX, POSEIDON2_AIR_REL_IDX,
26};
27
28#[derive(AlignedBytesBorrow, Clone)]
29#[repr(C)]
30struct DeferralOutputPrecompute {
31 rd_ptr: u32,
32 rs_ptr: u32,
33 deferral_idx: u32,
34}
35
36impl DeferralOutputExecutor {
37 #[inline(always)]
38 fn pre_compute_impl<F: PrimeField32>(
39 &self,
40 pc: u32,
41 inst: &Instruction<F>,
42 data: &mut DeferralOutputPrecompute,
43 ) -> Result<(), StaticProgramError> {
44 let Instruction {
45 opcode,
46 a,
47 b,
48 c,
49 d,
50 e,
51 ..
52 } = inst;
53
54 if opcode.local_opcode_idx(DeferralOpcode::CLASS_OFFSET) != DeferralOpcode::OUTPUT as usize
55 || d.as_canonical_u32() != RV32_REGISTER_AS
56 || e.as_canonical_u32() != RV32_MEMORY_AS
57 {
58 return Err(StaticProgramError::InvalidInstruction(pc));
59 }
60
61 *data = DeferralOutputPrecompute {
62 rd_ptr: a.as_canonical_u32(),
63 rs_ptr: b.as_canonical_u32(),
64 deferral_idx: c.as_canonical_u32(),
65 };
66 Ok(())
67 }
68}
69
70impl<F: PrimeField32> InterpreterExecutor<F> for DeferralOutputExecutor {
71 fn pre_compute_size(&self) -> usize {
72 size_of::<DeferralOutputPrecompute>()
73 }
74
75 #[cfg(not(feature = "tco"))]
76 fn pre_compute<Ctx>(
77 &self,
78 pc: u32,
79 inst: &Instruction<F>,
80 data: &mut [u8],
81 ) -> Result<ExecuteFunc<F, Ctx>, StaticProgramError>
82 where
83 Ctx: ExecutionCtxTrait,
84 {
85 let pre_compute: &mut DeferralOutputPrecompute = data.borrow_mut();
86 self.pre_compute_impl(pc, inst, pre_compute)?;
87 Ok(execute_e1_impl::<_, _>)
88 }
89
90 #[cfg(feature = "tco")]
91 fn handler<Ctx>(
92 &self,
93 pc: u32,
94 inst: &Instruction<F>,
95 data: &mut [u8],
96 ) -> Result<Handler<F, Ctx>, StaticProgramError>
97 where
98 Ctx: ExecutionCtxTrait,
99 {
100 let pre_compute: &mut DeferralOutputPrecompute = data.borrow_mut();
101 self.pre_compute_impl(pc, inst, pre_compute)?;
102 Ok(execute_e1_handler::<_, _>)
103 }
104}
105
106#[cfg(feature = "aot")]
107impl<F: PrimeField32> AotExecutor<F> for DeferralOutputExecutor {}
108
109impl<F: PrimeField32> InterpreterMeteredExecutor<F> for DeferralOutputExecutor {
110 fn metered_pre_compute_size(&self) -> usize {
111 size_of::<E2PreCompute<DeferralOutputPrecompute>>()
112 }
113
114 #[cfg(not(feature = "tco"))]
115 fn metered_pre_compute<Ctx>(
116 &self,
117 air_idx: usize,
118 pc: u32,
119 inst: &Instruction<F>,
120 data: &mut [u8],
121 ) -> Result<ExecuteFunc<F, Ctx>, StaticProgramError>
122 where
123 Ctx: MeteredExecutionCtxTrait,
124 {
125 let pre_compute: &mut E2PreCompute<DeferralOutputPrecompute> = data.borrow_mut();
126 pre_compute.chip_idx = air_idx as u32;
127 self.pre_compute_impl(pc, inst, &mut pre_compute.data)?;
128 Ok(execute_e2_impl::<_, _>)
129 }
130
131 #[cfg(feature = "tco")]
132 fn metered_handler<Ctx>(
133 &self,
134 air_idx: usize,
135 pc: u32,
136 inst: &Instruction<F>,
137 data: &mut [u8],
138 ) -> Result<Handler<F, Ctx>, StaticProgramError>
139 where
140 Ctx: MeteredExecutionCtxTrait,
141 {
142 let pre_compute: &mut E2PreCompute<DeferralOutputPrecompute> = data.borrow_mut();
143 pre_compute.chip_idx = air_idx as u32;
144 self.pre_compute_impl(pc, inst, &mut pre_compute.data)?;
145 Ok(execute_e2_handler::<_, _>)
146 }
147}
148
149#[cfg(feature = "aot")]
150impl<F: PrimeField32> AotMeteredExecutor<F> for DeferralOutputExecutor {}
151
152#[inline(always)]
153unsafe fn execute_e12_impl<F: PrimeField32, CTX: ExecutionCtxTrait>(
154 pre_compute: &DeferralOutputPrecompute,
155 exec_state: &mut VmExecState<F, GuestMemory, CTX>,
156) -> u32 {
157 let output_ptr = u32::from_le_bytes(exec_state.vm_read(RV32_REGISTER_AS, pre_compute.rd_ptr));
158 let input_ptr = u32::from_le_bytes(exec_state.vm_read(RV32_REGISTER_AS, pre_compute.rs_ptr));
159 let output_key_chunks: [[u8; DEFAULT_BLOCK_SIZE]; OUTPUT_TOTAL_MEMORY_OPS] = from_fn(|i| {
160 exec_state.vm_read(RV32_MEMORY_AS, input_ptr + (i * DEFAULT_BLOCK_SIZE) as u32)
161 });
162 let output_key: [u8; OUTPUT_TOTAL_BYTES] = join_memory_ops(output_key_chunks);
163 let (output_commit, output_len) = split_output(output_key);
164
165 let output_len_val = u64::from_le_bytes(output_len) as usize;
166
167 let num_rows = output_len_val / DIGEST_SIZE + 1;
170 debug_assert!(output_len_val.is_multiple_of(DIGEST_SIZE));
171
172 let output_raw = exec_state.streams.deferrals[pre_compute.deferral_idx as usize]
173 .get_output(&output_commit.to_vec())
174 .clone();
175 debug_assert_eq!(output_raw.len(), output_len_val);
176
177 for (row_idx, output_chunk) in output_raw.chunks_exact(DIGEST_SIZE).enumerate() {
178 let row_output_ptr = output_ptr + (row_idx * DIGEST_SIZE) as u32;
179 for chunk_idx in 0..DIGEST_MEMORY_OPS {
180 exec_state.vm_write::<u8, DEFAULT_BLOCK_SIZE>(
181 RV32_MEMORY_AS,
182 row_output_ptr + (chunk_idx * DEFAULT_BLOCK_SIZE) as u32,
183 &memory_op_chunk(output_chunk, chunk_idx),
184 );
185 }
186 }
187
188 let pc = exec_state.pc();
189 exec_state.set_pc(pc.wrapping_add(DEFAULT_PC_STEP));
190 num_rows as u32
191}
192
193#[create_handler]
194#[inline(always)]
195unsafe fn execute_e1_impl<F: PrimeField32, CTX: ExecutionCtxTrait>(
196 pre_compute: *const u8,
197 exec_state: &mut VmExecState<F, GuestMemory, CTX>,
198) {
199 let pre_compute: &DeferralOutputPrecompute =
200 from_raw_parts(pre_compute, size_of::<DeferralOutputPrecompute>()).borrow();
201 execute_e12_impl(pre_compute, exec_state);
202}
203
204#[create_handler]
205#[inline(always)]
206unsafe fn execute_e2_impl<F: PrimeField32, CTX: MeteredExecutionCtxTrait>(
207 pre_compute: *const u8,
208 exec_state: &mut VmExecState<F, GuestMemory, CTX>,
209) {
210 let pre_compute: &E2PreCompute<DeferralOutputPrecompute> = from_raw_parts(
211 pre_compute,
212 size_of::<E2PreCompute<DeferralOutputPrecompute>>(),
213 )
214 .borrow();
215 let height = execute_e12_impl(&pre_compute.data, exec_state);
216 exec_state
217 .ctx
218 .on_height_change(pre_compute.chip_idx as usize, height);
219
220 exec_state.ctx.on_height_change(
224 pre_compute.chip_idx as usize + (OUTPUT_AIR_REL_IDX - POSEIDON2_AIR_REL_IDX),
225 height,
226 );
227}