openvm_deferral_circuit/output/
execution.rs

1use 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    // Bytes are sponge-hashed and constrained against output_commit. The
168    // sponge rate is DIGEST_SIZE.
169    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    // The Poseidon2 peripheral chip's height also increases as a result of
221    // this opcode's execution. Computing an output commit from the raw output
222    // takes height Poseidon2 compressions.
223    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}