Skip to main content

ab_riscv_interpreter/v/zvexx/store/
zvexx_store_helpers.rs

1//! Opaque helpers for ZveXx extension
2
3use crate::v::vector_registers::VectorRegistersExt;
4use crate::v::zvexx::load::zvexx_load_helpers::{
5    check_register_group_alignment, mask_bit, read_group_element, snapshot_mask,
6};
7use crate::v::zvexx::zvexx_helpers::INSTRUCTION_SIZE;
8use crate::{ExecutionError, ProgramCounter, VirtualMemory, VirtualMemoryError};
9use ab_riscv_primitives::prelude::*;
10use core::fmt;
11use core::hint::cold_path;
12use core::num::NonZeroU8;
13
14/// Interpret `buf[..index_eew.bytes()]` as a little-endian unsigned integer and return it as
15/// `u64`. Used to convert a packed index element into a byte offset.
16///
17/// # Safety
18/// `index_eew.bytes() <= Eew::MAX_BYTES`, which is always true by construction.
19#[inline(always)]
20#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
21unsafe fn index_buf_to_u64(
22    buf: [u8; const { usize::from(Eew::MAX_BYTES) }],
23    index_eew: Eew,
24) -> u64 {
25    match index_eew {
26        Eew::E8 => u64::from(buf[0]),
27        Eew::E16 => u64::from(u16::from_le_bytes([buf[0], buf[1]])),
28        Eew::E32 => u64::from(u32::from_le_bytes([buf[0], buf[1], buf[2], buf[3]])),
29        Eew::E64 => u64::from_le_bytes(buf),
30    }
31}
32
33/// Write `eew`-sized data from `buf[..eew.bytes()]` to memory at `addr` (little-endian)
34#[inline(always)]
35#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
36fn write_mem_element(
37    memory: &mut impl VirtualMemory,
38    addr: u64,
39    eew: Eew,
40    buf: [u8; const { usize::from(Eew::MAX_BYTES) }],
41) -> Result<(), VirtualMemoryError> {
42    memory.write_slice(addr, &buf[..usize::from(eew.bytes_width())])
43}
44
45/// Validate a segment store's destination register group.
46///
47/// Like [`validate_segment_registers`] but omits the v0-overlap check, since
48/// segment stores read `vs3` as a source and the source/v0 overlap restriction
49/// applies only to load destinations.
50#[inline(always)]
51#[doc(hidden)]
52#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
53pub fn validate_segment_store_registers<Reg, Memory, PC, CustomError>(
54    program_counter: &PC,
55    vs3: VReg,
56    group_regs: NonZeroU8,
57    nf: Nf,
58) -> Result<(), ExecutionError<Reg::Type, CustomError>>
59where
60    Reg: Register,
61    PC: ProgramCounter<Reg::Type, Memory, CustomError>,
62{
63    if let Err(error) =
64        check_register_group_alignment::<Reg, _, _, _>(program_counter, vs3, group_regs)
65    {
66        cold_path();
67        return Err(error);
68    }
69    let total =
70        u32::from(vs3.to_bits()) + u32::from(nf.fields_per_segment()) * u32::from(group_regs.get());
71    if total > 32 {
72        cold_path();
73        return Err(ExecutionError::IllegalInstruction {
74            address: program_counter.old_pc(INSTRUCTION_SIZE),
75        });
76    }
77    Ok(())
78}
79
80/// Execute a unit-stride or unit-stride segment store.
81///
82/// Segment stride between elements is `nf * eew.bytes()`. Field `f` for element `i` is at
83/// `base + i * nf * eew.bytes() + f * eew.bytes()`. When `nf == 1` this degenerates to a
84/// plain unit-stride store.
85///
86/// # Safety
87/// - `vs3.to_bits() % group_regs == 0`
88/// - `vs3.to_bits() + nf * group_regs <= 32`
89/// - `vl <= group_regs * VLEN.bytes() / eew.bytes()` (all `vl` elements fit within the source
90///   register group; this holds when `vl` is the architectural `vl` and `group_regs` is the EMUL
91///   register count for the given `eew` and `vtype`)
92/// - When `vm=false`: `vs3` does not overlap `v0` (i.e. `vs3.to_bits() != 0`)
93#[inline(always)]
94#[expect(clippy::too_many_arguments, reason = "Internal API")]
95#[doc(hidden)]
96// TODO: #[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
97pub unsafe fn execute_unit_stride_store<Reg, ExtState, Memory, CustomError>(
98    ext_state: &mut ExtState,
99    memory: &mut Memory,
100    vs3: VReg,
101    vm: bool,
102    base: u64,
103    eew: Eew,
104    group_regs: NonZeroU8,
105    nf: Nf,
106) -> Result<(), ExecutionError<Reg::Type, CustomError>>
107where
108    Reg: Register,
109    ExtState: VectorRegistersExt<Reg, CustomError>,
110    [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
111    Memory: VirtualMemory,
112    CustomError: fmt::Debug,
113{
114    let group_regs = group_regs.get();
115    let vl = ext_state.vl();
116    let vstart = ext_state.vstart();
117    let elem_bytes = eew.bytes_width();
118    let segment_stride = u64::from(nf.fields_per_segment() * elem_bytes);
119    // SAFETY: `vl <= VLMAX <= VLEN`
120    let mask_buf = unsafe { snapshot_mask(ext_state.read_vregs(), vm, vl) };
121    for i in vstart.range_to(vl) {
122        if !vm && !mask_bit(&mask_buf, i) {
123            continue;
124        }
125        let elem_base = base.wrapping_add(u64::from(i) * segment_stride);
126        for f in 0..nf.fields_per_segment() {
127            let addr = elem_base.wrapping_add(u64::from(f * elem_bytes));
128            // SAFETY: Guaranteed by function contract
129            let field_base_reg =
130                unsafe { VReg::from_bits(vs3.to_bits() + f * group_regs).unwrap_unchecked() };
131            // SAFETY: need `field_base_reg + i / (VLEN.bytes() / elem_bytes) < 32`.
132            //
133            // Let `elems_per_reg = VLEN.bytes() / elem_bytes`.
134            // `i < vl <= group_regs * elems_per_reg` (precondition), so
135            // `i / elems_per_reg < group_regs`.
136            //
137            // `field_base_reg = vs3.to_bits() + f * group_regs`. Since `f < nf` and the
138            // precondition guarantees `vs3.to_bits() + nf * group_regs <= 32`:
139            // `field_base_reg + group_regs <= vs3.to_bits() + (f+1) * group_regs
140            //                             <= vs3.to_bits() + nf * group_regs <= 32`.
141            //
142            // Therefore,
143            // `field_base_reg + i / elems_per_reg < field_base_reg + group_regs <= 32`.
144            let data =
145                unsafe { read_group_element(ext_state.read_vregs(), field_base_reg, i, eew) };
146            // Record the current element index in `vstart` so that, on a memory fault, the failing
147            // element can be identified and the operation can be restarted
148            if let Err(error) = write_mem_element(memory, addr, eew, data) {
149                cold_path();
150                ext_state.set_vstart(Vstart::from(i));
151                return Err(ExecutionError::MemoryAccess(error));
152            }
153        }
154    }
155    ext_state.reset_vstart();
156    Ok(())
157}
158
159/// Execute a strided or strided-segment store.
160///
161/// The address of element `i`, field `f` is:
162///   `base.wrapping_add(i.wrapping_mul(stride) as u64).wrapping_add(f * eew.bytes())`
163///
164/// `stride` is the raw XLEN register value reinterpreted as a signed integer, matching the RVV
165/// specification where the stride operand is a two's-complement signed offset.
166///
167/// # Safety
168/// - `vs3.to_bits() % group_regs == 0`
169/// - `vs3.to_bits() + nf * group_regs <= 32`
170/// - `vl <= group_regs * VLEN.bytes() / eew.bytes()`
171/// - When `vm=false`: `vs3.to_bits() != 0`
172#[inline(always)]
173#[expect(clippy::too_many_arguments, reason = "Internal API")]
174#[doc(hidden)]
175// TODO: #[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
176pub unsafe fn execute_strided_store<Reg, ExtState, Memory, CustomError>(
177    ext_state: &mut ExtState,
178    memory: &mut Memory,
179    vs3: VReg,
180    vm: bool,
181    base: u64,
182    stride: i64,
183    eew: Eew,
184    group_regs: NonZeroU8,
185    nf: Nf,
186) -> Result<(), ExecutionError<Reg::Type, CustomError>>
187where
188    Reg: Register,
189    ExtState: VectorRegistersExt<Reg, CustomError>,
190    [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
191    Memory: VirtualMemory,
192    CustomError: fmt::Debug,
193{
194    let group_regs = group_regs.get();
195    let vl = ext_state.vl();
196    let vstart = ext_state.vstart();
197    let elem_bytes = eew.bytes_width();
198    // SAFETY: `vl <= VLMAX <= VLEN`
199    let mask_buf = unsafe { snapshot_mask(ext_state.read_vregs(), vm, vl) };
200    for i in vstart.range_to(vl) {
201        if !vm && !mask_bit(&mask_buf, i) {
202            continue;
203        }
204        let elem_base = base.wrapping_add(i64::from(i).wrapping_mul(stride).cast_unsigned());
205        for f in 0..nf.fields_per_segment() {
206            let addr = elem_base.wrapping_add(u64::from(f * elem_bytes));
207            // SAFETY: Guaranteed by function contract
208            let field_base_reg =
209                unsafe { VReg::from_bits(vs3.to_bits() + f * group_regs).unwrap_unchecked() };
210            // SAFETY: same argument as `execute_unit_stride_store`; `field_base_reg +
211            // i / elems_per_reg < field_base_reg + group_regs <= vs3.to_bits() + nf *
212            // group_regs <= 32`.
213            let data =
214                unsafe { read_group_element(ext_state.read_vregs(), field_base_reg, i, eew) };
215            // Record the current element index in `vstart` so that, on a memory fault, the failing
216            // element can be identified and the operation can be restarted
217            if let Err(error) = write_mem_element(memory, addr, eew, data) {
218                cold_path();
219                ext_state.set_vstart(Vstart::from(i));
220                return Err(ExecutionError::MemoryAccess(error));
221            }
222        }
223    }
224    ext_state.reset_vstart();
225    Ok(())
226}
227
228/// Execute an indexed (unordered or ordered) store or indexed-segment store.
229///
230/// The effective address of element `i`, field `f` is:
231///   `base + index[i] + f * eew.bytes()`
232/// where `index[i]` is element `i` of the index register group `vs2`, interpreted as an
233/// unsigned integer of width `index_eew`.
234///
235/// `data_eew` is the element width of the data being stored (from `vtype.vsew`).
236/// `index_eew` is the element width of the indices (from the instruction encoding).
237///
238/// # Safety
239/// - `vs3.to_bits() % data_group_regs == 0`
240/// - `vs3.to_bits() + nf * data_group_regs <= 32`
241/// - `vs2` register group is aligned and fits within `[0, 32)` (caller must verify via
242///   `check_register_group_alignment` before calling)
243/// - `vl <= data_group_regs * VLEN.bytes() / data_eew.bytes()`
244/// - `vl <= index_group_regs * VLEN.bytes() / index_eew.bytes()` (caller must verify)
245/// - When `vm=false`: `vs3.to_bits() != 0`
246#[inline(always)]
247#[expect(clippy::too_many_arguments, reason = "Internal API")]
248#[doc(hidden)]
249// TODO: #[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
250pub unsafe fn execute_indexed_store<Reg, ExtState, Memory, CustomError>(
251    ext_state: &mut ExtState,
252    memory: &mut Memory,
253    vs3: VReg,
254    vs2: VReg,
255    vm: bool,
256    base: u64,
257    data_eew: Eew,
258    index_eew: Eew,
259    data_group_regs: NonZeroU8,
260    nf: Nf,
261) -> Result<(), ExecutionError<Reg::Type, CustomError>>
262where
263    Reg: Register,
264    ExtState: VectorRegistersExt<Reg, CustomError>,
265    [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
266    Memory: VirtualMemory,
267    CustomError: fmt::Debug,
268{
269    let data_group_regs = data_group_regs.get();
270    let vl = ext_state.vl();
271    let vstart = ext_state.vstart();
272    let data_elem_bytes = data_eew.bytes_width();
273    // SAFETY: `vl <= VLMAX <= VLEN`
274    let mask_buf = unsafe { snapshot_mask(ext_state.read_vregs(), vm, vl) };
275    for i in vstart.range_to(vl) {
276        if !vm && !mask_bit(&mask_buf, i) {
277            continue;
278        }
279        // SAFETY: `i < vl <= index_group_regs * VLEN.bytes() / index_eew.bytes()` (precondition),
280        // so `vs2.to_bits() + i / (VLEN.bytes() / index_eew.bytes()) <
281        //     vs2.to_bits() + index_group_regs <= 32`
282        let index_buf = unsafe { read_group_element(ext_state.read_vregs(), vs2, i, index_eew) };
283        // SAFETY: `index_eew.bytes() <= Eew::MAX_BYTES` always holds.
284        let offset = unsafe { index_buf_to_u64(index_buf, index_eew) };
285        let elem_base = base.wrapping_add(offset);
286        for f in 0..nf.fields_per_segment() {
287            let addr = elem_base.wrapping_add(u64::from(f) * u64::from(data_elem_bytes));
288            // SAFETY: Guaranteed by function contract
289            let field_base_reg =
290                unsafe { VReg::from_bits(vs3.to_bits() + f * data_group_regs).unwrap_unchecked() };
291            // SAFETY: `i < vl <= data_group_regs * VLEN.bytes() / data_eew.bytes()` (precondition),
292            // so `field_base_reg + i / elems_per_reg < field_base_reg + data_group_regs
293            //                                    <= vs3.to_bits() + nf * data_group_regs <= 32`.
294            let data =
295                unsafe { read_group_element(ext_state.read_vregs(), field_base_reg, i, data_eew) };
296            // Record the current element index in `vstart` so that, on a memory fault, the failing
297            // element can be identified and the operation can be restarted
298            if let Err(error) = write_mem_element(memory, addr, data_eew, data) {
299                cold_path();
300                ext_state.set_vstart(Vstart::from(i));
301                return Err(ExecutionError::MemoryAccess(error));
302            }
303        }
304    }
305    ext_state.reset_vstart();
306    Ok(())
307}