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::{VectorRegisterFile, VectorRegistersExt};
4use crate::v::zvexx::load::zvexx_load_helpers::{
5    check_register_group_alignment, mask_bit, 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;
11
12/// Interpret `buf[..index_eew.bytes()]` as a little-endian unsigned integer and return it as
13/// `u64`. Used to convert a packed index element into a byte offset.
14///
15/// # Safety
16/// `index_eew.bytes() <= Eew::MAX_BYTES`, which is always true by construction.
17#[inline(always)]
18#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
19unsafe fn index_buf_to_u64(
20    buf: [u8; const { usize::from(Eew::MAX_BYTES) }],
21    index_eew: Eew,
22) -> u64 {
23    match index_eew {
24        Eew::E8 => u64::from(buf[0]),
25        Eew::E16 => u64::from(u16::from_le_bytes([buf[0], buf[1]])),
26        Eew::E32 => u64::from(u32::from_le_bytes([buf[0], buf[1], buf[2], buf[3]])),
27        Eew::E64 => u64::from_le_bytes(buf),
28    }
29}
30
31/// Write `eew`-sized data from `buf[..eew.bytes()]` to memory at `addr` (little-endian)
32#[inline(always)]
33#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
34fn write_mem_element(
35    memory: &mut impl VirtualMemory,
36    addr: u64,
37    eew: Eew,
38    buf: [u8; const { usize::from(Eew::MAX_BYTES) }],
39) -> Result<(), VirtualMemoryError> {
40    memory.write_slice(addr, &buf[..usize::from(eew.bytes_width())])
41}
42
43/// Validate a segment store's destination register group.
44///
45/// Like [`validate_segment_registers`] but omits the v0-overlap check, since
46/// segment stores read `vs3` as a source and the source/v0 overlap restriction
47/// applies only to load destinations.
48#[inline(always)]
49#[doc(hidden)]
50#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
51pub fn validate_segment_store_registers<Reg, Memory, PC>(
52    program_counter: &PC,
53    vs3: VReg,
54    group_regs: VRegGroupSize,
55    nf: Nf,
56) -> Result<(), ExecutionError<Reg::Type>>
57where
58    Reg: Register,
59    PC: ProgramCounter<Reg::Type, Memory>,
60{
61    if let Err(error) =
62        check_register_group_alignment::<Reg, _, _>(program_counter, vs3, group_regs)
63    {
64        cold_path();
65        return Err(error);
66    }
67    let nf_group_regs = u32::from(nf.fields_per_segment()) * u32::from(group_regs.get());
68    let total = u32::from(vs3.to_bits()) + nf_group_regs;
69    // Per spec, `NFIELDS * EMUL` must not exceed 8 for segment loads/stores, regardless of whether
70    // the field groups would otherwise fit within the 32 vector registers
71    if nf_group_regs > 8 || total > 32 {
72        cold_path();
73        return Err(ExecutionError::IllegalInstruction {
74            address: PackedAddress::new(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#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
97pub unsafe fn execute_unit_stride_store<Reg, Env, Memory>(
98    env: &mut Env,
99    memory: &mut Memory,
100    vs3: VReg,
101    vm: bool,
102    base: u64,
103    eew: Eew,
104    group_regs: VRegGroupSize,
105    nf: Nf,
106) -> Result<(), ExecutionError<Reg::Type>>
107where
108    Reg: Register,
109    Env: VectorRegistersExt<Reg>,
110    [(); SUPPORTED_ELEN_VLEN::<{ Env::ELEN }, { Env::VLEN }>]:,
111    Memory: VirtualMemory,
112{
113    let group_regs = group_regs.get();
114    let vl = env.vl();
115    let vstart = env.vstart();
116    let elem_bytes = eew.bytes_width();
117
118    let range = vstart.range_to(vl);
119
120    // Unmasked non-segment store is a plain copy of the contiguous element range of the register
121    // group into a contiguous memory range, as long as the whole range is writable. If it is not,
122    // the element-wise path below is what determines the faulting element and writes everything
123    // before it, which is the same outcome a partially written copy would have.
124    if vm && nf.fields_per_segment() == 1 && !range.is_empty() {
125        let first = *range.start();
126        let len = range.len() * usize::from(elem_bytes);
127        let addr = base.wrapping_add(u64::from(first) * u64::from(elem_bytes));
128        let offset = VectorRegisterFile::<{ Env::VLEN }>::element_offset(vs3, first, eew);
129        // SAFETY: Elements `vstart..vl` all lie within the register group, which ends within the
130        // register file (precondition), so `offset + len <= 32 * VLEN.bytes()`
131        let data = unsafe {
132            env.read_vregs()
133                .as_bytes()
134                .as_flattened()
135                .get_unchecked(offset..offset + len)
136        };
137        if memory.write_slice(addr, data).is_ok() {
138            env.reset_vstart();
139            return Ok(());
140        }
141        cold_path();
142    }
143
144    let segment_stride = u64::from(nf.fields_per_segment() * elem_bytes);
145    // SAFETY: `vl <= VLMAX <= VLEN`
146    let mask_buf = unsafe { snapshot_mask(env.read_vregs(), vm, vl) };
147    for i in range {
148        if !vm && !mask_bit(&mask_buf, i) {
149            continue;
150        }
151        let elem_base = base.wrapping_add(u64::from(i) * segment_stride);
152        for f in 0..nf.fields_per_segment() {
153            let addr = elem_base.wrapping_add(u64::from(f * elem_bytes));
154            // SAFETY: Guaranteed by function contract
155            let field_base_reg =
156                unsafe { VReg::from_bits(vs3.to_bits() + f * group_regs).unwrap_unchecked() };
157            // SAFETY: need `field_base_reg + i / (VLEN.bytes() / elem_bytes) < 32`.
158            //
159            // Let `elems_per_reg = VLEN.bytes() / elem_bytes`.
160            // `i < vl <= group_regs * elems_per_reg` (precondition), so
161            // `i / elems_per_reg < group_regs`.
162            //
163            // `field_base_reg = vs3.to_bits() + f * group_regs`. Since `f < nf` and the
164            // precondition guarantees `vs3.to_bits() + nf * group_regs <= 32`:
165            // `field_base_reg + group_regs <= vs3.to_bits() + (f+1) * group_regs
166            //                             <= vs3.to_bits() + nf * group_regs <= 32`.
167            //
168            // Therefore,
169            // `field_base_reg + i / elems_per_reg < field_base_reg + group_regs <= 32`.
170            let data = unsafe {
171                env.read_vregs()
172                    .read_element(field_base_reg, i, eew)
173                    .to_le_bytes()
174            };
175            // Record the current element index in `vstart` so that, on a memory fault, the failing
176            // element can be identified and the operation can be restarted
177            if let Err(error) = write_mem_element(memory, addr, eew, data) {
178                cold_path();
179                env.set_vstart(Vstart::from(i));
180                return Err(ExecutionError::from(error));
181            }
182        }
183    }
184    env.reset_vstart();
185    Ok(())
186}
187
188/// Execute a strided or strided-segment store.
189///
190/// The address of element `i`, field `f` is:
191///   `base.wrapping_add(i.wrapping_mul(stride) as u64).wrapping_add(f * eew.bytes())`
192///
193/// `stride` is the raw XLEN register value reinterpreted as a signed integer, matching the RVV
194/// specification where the stride operand is a two's-complement signed offset.
195///
196/// # Safety
197/// - `vs3.to_bits() % group_regs == 0`
198/// - `vs3.to_bits() + nf * group_regs <= 32`
199/// - `vl <= group_regs * VLEN.bytes() / eew.bytes()`
200/// - When `vm=false`: `vs3.to_bits() != 0`
201#[inline(always)]
202#[expect(clippy::too_many_arguments, reason = "Internal API")]
203#[doc(hidden)]
204#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
205pub unsafe fn execute_strided_store<Reg, Env, Memory>(
206    env: &mut Env,
207    memory: &mut Memory,
208    vs3: VReg,
209    vm: bool,
210    base: u64,
211    stride: i64,
212    eew: Eew,
213    group_regs: VRegGroupSize,
214    nf: Nf,
215) -> Result<(), ExecutionError<Reg::Type>>
216where
217    Reg: Register,
218    Env: VectorRegistersExt<Reg>,
219    [(); SUPPORTED_ELEN_VLEN::<{ Env::ELEN }, { Env::VLEN }>]:,
220    Memory: VirtualMemory,
221{
222    let group_regs = group_regs.get();
223    let vl = env.vl();
224    let vstart = env.vstart();
225    let elem_bytes = eew.bytes_width();
226    // SAFETY: `vl <= VLMAX <= VLEN`
227    let mask_buf = unsafe { snapshot_mask(env.read_vregs(), vm, vl) };
228    for i in vstart.range_to(vl) {
229        if !vm && !mask_bit(&mask_buf, i) {
230            continue;
231        }
232        let elem_base = base.wrapping_add(i64::from(i).wrapping_mul(stride).cast_unsigned());
233        for f in 0..nf.fields_per_segment() {
234            let addr = elem_base.wrapping_add(u64::from(f * elem_bytes));
235            // SAFETY: Guaranteed by function contract
236            let field_base_reg =
237                unsafe { VReg::from_bits(vs3.to_bits() + f * group_regs).unwrap_unchecked() };
238            // SAFETY: same argument as `execute_unit_stride_store`; `field_base_reg +
239            // i / elems_per_reg < field_base_reg + group_regs <= vs3.to_bits() + nf *
240            // group_regs <= 32`.
241            let data = unsafe {
242                env.read_vregs()
243                    .read_element(field_base_reg, i, eew)
244                    .to_le_bytes()
245            };
246            // Record the current element index in `vstart` so that, on a memory fault, the failing
247            // element can be identified and the operation can be restarted
248            if let Err(error) = write_mem_element(memory, addr, eew, data) {
249                cold_path();
250                env.set_vstart(Vstart::from(i));
251                return Err(ExecutionError::from(error));
252            }
253        }
254    }
255    env.reset_vstart();
256    Ok(())
257}
258
259/// Execute an indexed (unordered or ordered) store or indexed-segment store.
260///
261/// The effective address of element `i`, field `f` is:
262///   `base + index[i] + f * eew.bytes()`
263/// where `index[i]` is element `i` of the index register group `vs2`, interpreted as an
264/// unsigned integer of width `index_eew`.
265///
266/// `data_eew` is the element width of the data being stored (from `vtype.vsew`).
267/// `index_eew` is the element width of the indices (from the instruction encoding).
268///
269/// # Safety
270/// - `vs3.to_bits() % data_group_regs == 0`
271/// - `vs3.to_bits() + nf * data_group_regs <= 32`
272/// - `vs2` register group is aligned and fits within `[0, 32)` (caller must verify via
273///   `check_register_group_alignment` before calling)
274/// - `vl <= data_group_regs * VLEN.bytes() / data_eew.bytes()`
275/// - `vl <= index_group_regs * VLEN.bytes() / index_eew.bytes()` (caller must verify)
276/// - When `vm=false`: `vs3.to_bits() != 0`
277#[inline(always)]
278#[expect(clippy::too_many_arguments, reason = "Internal API")]
279#[doc(hidden)]
280#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
281pub unsafe fn execute_indexed_store<Reg, Env, Memory>(
282    env: &mut Env,
283    memory: &mut Memory,
284    vs3: VReg,
285    vs2: VReg,
286    vm: bool,
287    base: u64,
288    data_eew: Eew,
289    index_eew: Eew,
290    data_group_regs: VRegGroupSize,
291    nf: Nf,
292) -> Result<(), ExecutionError<Reg::Type>>
293where
294    Reg: Register,
295    Env: VectorRegistersExt<Reg>,
296    [(); SUPPORTED_ELEN_VLEN::<{ Env::ELEN }, { Env::VLEN }>]:,
297    Memory: VirtualMemory,
298{
299    let data_group_regs = data_group_regs.get();
300    let vl = env.vl();
301    let vstart = env.vstart();
302    let data_elem_bytes = data_eew.bytes_width();
303    // SAFETY: `vl <= VLMAX <= VLEN`
304    let mask_buf = unsafe { snapshot_mask(env.read_vregs(), vm, vl) };
305    for i in vstart.range_to(vl) {
306        if !vm && !mask_bit(&mask_buf, i) {
307            continue;
308        }
309        // SAFETY: `i < vl <= index_group_regs * VLEN.bytes() / index_eew.bytes()` (precondition),
310        // so `vs2.to_bits() + i / (VLEN.bytes() / index_eew.bytes()) <
311        //     vs2.to_bits() + index_group_regs <= 32`
312        let index_buf = unsafe {
313            env.read_vregs()
314                .read_element(vs2, i, index_eew)
315                .to_le_bytes()
316        };
317        // SAFETY: `index_eew.bytes() <= Eew::MAX_BYTES` always holds.
318        let offset = unsafe { index_buf_to_u64(index_buf, index_eew) };
319        let elem_base = base.wrapping_add(offset);
320        for f in 0..nf.fields_per_segment() {
321            let addr = elem_base.wrapping_add(u64::from(f) * u64::from(data_elem_bytes));
322            // SAFETY: Guaranteed by function contract
323            let field_base_reg =
324                unsafe { VReg::from_bits(vs3.to_bits() + f * data_group_regs).unwrap_unchecked() };
325            // SAFETY: `i < vl <= data_group_regs * VLEN.bytes() / data_eew.bytes()` (precondition),
326            // so `field_base_reg + i / elems_per_reg < field_base_reg + data_group_regs
327            //                                    <= vs3.to_bits() + nf * data_group_regs <= 32`.
328            let data = unsafe {
329                env.read_vregs()
330                    .read_element(field_base_reg, i, data_eew)
331                    .to_le_bytes()
332            };
333            // Record the current element index in `vstart` so that, on a memory fault, the failing
334            // element can be identified and the operation can be restarted
335            if let Err(error) = write_mem_element(memory, addr, data_eew, data) {
336                cold_path();
337                env.set_vstart(Vstart::from(i));
338                return Err(ExecutionError::from(error));
339            }
340        }
341    }
342    env.reset_vstart();
343    Ok(())
344}