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