openvm_keccak256_circuit/keccakf_op/
execution.rs

1use std::{
2    borrow::{Borrow, BorrowMut},
3    convert::TryInto,
4    mem::size_of,
5};
6
7use openvm_circuit::{
8    arch::{StaticProgramError, *},
9    system::memory::online::GuestMemory,
10};
11use openvm_circuit_primitives_derive::AlignedBytesBorrow;
12use openvm_instructions::{
13    instruction::Instruction,
14    program::DEFAULT_PC_STEP,
15    riscv::{RV32_MEMORY_AS, RV32_REGISTER_AS},
16};
17use openvm_stark_backend::p3_field::PrimeField32;
18use p3_keccak_air::NUM_ROUNDS;
19
20use super::{KeccakfExecutor, NUM_OP_ROWS_PER_INS};
21use crate::{keccakf_op::keccakf_postimage_bytes, KECCAK_WIDTH_BYTES, KECCAK_WORD_SIZE};
22
23#[derive(AlignedBytesBorrow, Clone)]
24#[repr(C)]
25struct KeccakfPreCompute {
26    a: u8,
27}
28
29impl KeccakfExecutor {
30    fn pre_compute_impl<F: PrimeField32>(
31        &self,
32        pc: u32,
33        inst: &Instruction<F>,
34        data: &mut KeccakfPreCompute,
35    ) -> Result<(), StaticProgramError> {
36        let Instruction {
37            opcode: _,
38            a,
39            b: _,
40            c: _,
41            d,
42            e,
43            ..
44        } = inst;
45
46        let e_u32 = e.as_canonical_u32();
47        if d.as_canonical_u32() != RV32_REGISTER_AS || e_u32 != RV32_MEMORY_AS {
48            return Err(StaticProgramError::InvalidInstruction(pc));
49        }
50
51        *data = KeccakfPreCompute {
52            a: a.as_canonical_u32() as u8,
53        };
54
55        Ok(())
56    }
57}
58
59impl<F: PrimeField32> InterpreterExecutor<F> for KeccakfExecutor {
60    fn pre_compute_size(&self) -> usize {
61        size_of::<KeccakfPreCompute>()
62    }
63
64    #[cfg(not(feature = "tco"))]
65    fn pre_compute<Ctx>(
66        &self,
67        pc: u32,
68        inst: &Instruction<F>,
69        data: &mut [u8],
70    ) -> Result<ExecuteFunc<F, Ctx>, StaticProgramError>
71    where
72        Ctx: ExecutionCtxTrait,
73    {
74        let data: &mut KeccakfPreCompute = data.borrow_mut();
75        self.pre_compute_impl(pc, inst, data)?;
76        Ok(execute_e1_impl::<_, _>)
77    }
78
79    #[cfg(feature = "tco")]
80    fn handler<Ctx>(
81        &self,
82        pc: u32,
83        inst: &Instruction<F>,
84        data: &mut [u8],
85    ) -> Result<Handler<F, Ctx>, StaticProgramError>
86    where
87        Ctx: ExecutionCtxTrait,
88    {
89        let data: &mut KeccakfPreCompute = data.borrow_mut();
90        self.pre_compute_impl(pc, inst, data)?;
91        Ok(execute_e1_handler)
92    }
93}
94
95#[cfg(feature = "aot")]
96impl<F: PrimeField32> AotExecutor<F> for KeccakfExecutor {}
97
98impl<F: PrimeField32> InterpreterMeteredExecutor<F> for KeccakfExecutor {
99    fn metered_pre_compute_size(&self) -> usize {
100        size_of::<E2PreCompute<KeccakfPreCompute>>()
101    }
102
103    #[cfg(not(feature = "tco"))]
104    fn metered_pre_compute<Ctx>(
105        &self,
106        chip_idx: usize,
107        pc: u32,
108        inst: &Instruction<F>,
109        data: &mut [u8],
110    ) -> Result<ExecuteFunc<F, Ctx>, StaticProgramError>
111    where
112        Ctx: MeteredExecutionCtxTrait,
113    {
114        let data: &mut E2PreCompute<KeccakfPreCompute> = data.borrow_mut();
115        data.chip_idx = chip_idx as u32;
116        self.pre_compute_impl(pc, inst, &mut data.data)?;
117        Ok(execute_e2_impl::<_, _>)
118    }
119
120    #[cfg(feature = "tco")]
121    fn metered_handler<Ctx>(
122        &self,
123        chip_idx: usize,
124        pc: u32,
125        inst: &Instruction<F>,
126        data: &mut [u8],
127    ) -> Result<Handler<F, Ctx>, StaticProgramError>
128    where
129        Ctx: MeteredExecutionCtxTrait,
130    {
131        let data: &mut E2PreCompute<KeccakfPreCompute> = data.borrow_mut();
132        data.chip_idx = chip_idx as u32;
133        self.pre_compute_impl(pc, inst, &mut data.data)?;
134        Ok(execute_e2_handler)
135    }
136}
137
138#[cfg(feature = "aot")]
139impl<F: PrimeField32> AotMeteredExecutor<F> for KeccakfExecutor {}
140
141#[create_handler]
142#[inline(always)]
143unsafe fn execute_e1_impl<F: PrimeField32, CTX: ExecutionCtxTrait>(
144    pre_compute: *const u8,
145    exec_state: &mut VmExecState<F, GuestMemory, CTX>,
146) {
147    let pre_compute: &KeccakfPreCompute =
148        std::slice::from_raw_parts(pre_compute, size_of::<KeccakfPreCompute>()).borrow();
149    execute_e12_impl::<F, CTX, true>(pre_compute, exec_state);
150}
151
152#[inline(always)]
153unsafe fn execute_e12_impl<F: PrimeField32, CTX: ExecutionCtxTrait, const IS_E1: bool>(
154    pre_compute: &KeccakfPreCompute,
155    exec_state: &mut VmExecState<F, GuestMemory, CTX>,
156) {
157    let rd_ptr = pre_compute.a as u32;
158    let buffer_ptr_limbs: [u8; 4] = exec_state.vm_read(RV32_REGISTER_AS, rd_ptr);
159    let buffer_ptr = u32::from_le_bytes(buffer_ptr_limbs);
160
161    let preimage: &[u8] =
162        exec_state.host_read_slice(RV32_MEMORY_AS, buffer_ptr, KECCAK_WIDTH_BYTES);
163    let postimage = keccakf_postimage_bytes(preimage.try_into().unwrap());
164
165    if IS_E1 {
166        exec_state.vm_write(RV32_MEMORY_AS, buffer_ptr, &postimage);
167    } else {
168        for (word_idx, word) in postimage.chunks_exact(KECCAK_WORD_SIZE).enumerate() {
169            exec_state.vm_write::<u8, KECCAK_WORD_SIZE>(
170                RV32_MEMORY_AS,
171                buffer_ptr + (word_idx * KECCAK_WORD_SIZE) as u32,
172                word.try_into().unwrap(),
173            );
174        }
175    }
176
177    let pc = exec_state.pc();
178    exec_state.set_pc(pc.wrapping_add(DEFAULT_PC_STEP));
179}
180
181#[create_handler]
182#[inline(always)]
183unsafe fn execute_e2_impl<F: PrimeField32, CTX: MeteredExecutionCtxTrait>(
184    pre_compute: *const u8,
185    exec_state: &mut VmExecState<F, GuestMemory, CTX>,
186) {
187    let pre_compute: &E2PreCompute<KeccakfPreCompute> =
188        std::slice::from_raw_parts(pre_compute, size_of::<E2PreCompute<KeccakfPreCompute>>())
189            .borrow();
190
191    let op_air_idx = pre_compute.chip_idx as usize;
192
193    // Update KeccakfOpChip height (2 rows per instruction)
194    exec_state
195        .ctx
196        .on_height_change(op_air_idx, NUM_OP_ROWS_PER_INS as u32);
197
198    // HACK: KeccakfPermAir is added right before KeccakfOpAir in extend_circuit,
199    // and due to reverse ordering of AIR indices, perm_air_idx = op_air_idx + 1.
200    // See extension/mod.rs extend_circuit for the ordering.
201    let perm_air_idx = op_air_idx + 1;
202
203    // Update KeccakfPermChip height (24 rows per keccakf permutation)
204    exec_state
205        .ctx
206        .on_height_change(perm_air_idx, NUM_ROUNDS as u32);
207
208    execute_e12_impl::<F, CTX, false>(&pre_compute.data, exec_state);
209}