use crate::account_view::AccountView;
use crate::borrow::Ref;
use crate::error::ProgramError;
pub unsafe trait Projectable: Copy + 'static {}
unsafe impl Projectable for u8 {}
unsafe impl Projectable for u16 {}
unsafe impl Projectable for u32 {}
unsafe impl Projectable for u64 {}
unsafe impl Projectable for u128 {}
unsafe impl Projectable for i8 {}
unsafe impl Projectable for i16 {}
unsafe impl Projectable for i32 {}
unsafe impl Projectable for i64 {}
unsafe impl Projectable for i128 {}
unsafe impl Projectable for [u8; 32] {}
unsafe impl Projectable for [u8; 64] {}
pub unsafe trait SafeProjectable: Projectable {}
unsafe impl<T: Projectable> SafeProjectable for T where Self: private::NonZeroSized {}
mod private {
pub trait NonZeroSized {}
impl<T: Copy + 'static> NonZeroSized for T {}
}
#[inline]
pub fn project_safe<'a, T: SafeProjectable>(
account: &'a AccountView<'a>,
offset: usize,
expected_disc: Option<u8>,
) -> Result<Ref<'a, T>, ProgramError> {
const {
assert!(
core::mem::size_of::<T>() > 0,
"project_safe: T must be non-zero-sized"
);
}
project::<T>(account, offset, expected_disc)
}
#[inline]
pub unsafe fn project_safe_mut<'a, T: SafeProjectable>(
account: &'a AccountView<'a>,
offset: usize,
expected_disc: Option<u8>,
) -> Result<&'a mut T, ProgramError> {
const {
assert!(
core::mem::size_of::<T>() > 0,
"project_safe_mut: T must be non-zero-sized"
);
}
unsafe { project_mut::<T>(account, offset, expected_disc) }
}
#[inline]
pub fn project<'a, T: Projectable>(
account: &'a AccountView<'a>,
offset: usize,
expected_disc: Option<u8>,
) -> Result<Ref<'a, T>, ProgramError> {
let data_len = account.data_len();
let type_size = core::mem::size_of::<T>();
if offset
.checked_add(type_size)
.is_none_or(|end| end > data_len)
{
return Err(ProgramError::AccountDataTooSmall);
}
if let Some(disc) = expected_disc {
if account.disc() != disc {
return Err(ProgramError::InvalidAccountData);
}
}
let data_ptr = account.data_ptr_unchecked();
let target_ptr = unsafe { data_ptr.add(offset) };
let align = core::mem::align_of::<T>();
if !(target_ptr as usize).is_multiple_of(align) {
return Err(ProgramError::InvalidAccountData);
}
let state_ptr = account.acquire_shared()?;
Ok(Ref::new(unsafe { &*(target_ptr as *const T) }, state_ptr))
}
#[inline]
pub unsafe fn project_mut<'a, T: Projectable>(
account: &'a AccountView<'a>,
offset: usize,
expected_disc: Option<u8>,
) -> Result<&'a mut T, ProgramError> {
let data_len = account.data_len();
let type_size = core::mem::size_of::<T>();
if offset
.checked_add(type_size)
.is_none_or(|end| end > data_len)
{
return Err(ProgramError::AccountDataTooSmall);
}
if let Some(disc) = expected_disc {
if account.disc() != disc {
return Err(ProgramError::InvalidAccountData);
}
}
let data_ptr = account.data_ptr_unchecked();
let target_ptr = unsafe { data_ptr.add(offset) };
let align = core::mem::align_of::<T>();
if !(target_ptr as usize).is_multiple_of(align) {
return Err(ProgramError::InvalidAccountData);
}
Ok(unsafe { &mut *(target_ptr as *mut T) })
}
#[inline]
pub fn project_slice<'a, T: Projectable>(
account: &'a AccountView<'a>,
offset: usize,
count: usize,
) -> Result<Ref<'a, [T]>, ProgramError> {
let data_len = account.data_len();
let type_size = core::mem::size_of::<T>();
let total = count
.checked_mul(type_size)
.ok_or(ProgramError::ArithmeticOverflow)?;
if offset.checked_add(total).is_none_or(|end| end > data_len) {
return Err(ProgramError::AccountDataTooSmall);
}
let data_ptr = account.data_ptr_unchecked();
let target_ptr = unsafe { data_ptr.add(offset) };
let align = core::mem::align_of::<T>();
if !(target_ptr as usize).is_multiple_of(align) {
return Err(ProgramError::InvalidAccountData);
}
let state_ptr = account.acquire_shared()?;
Ok(Ref::new(
unsafe { core::slice::from_raw_parts(target_ptr as *const T, count) },
state_ptr,
))
}
#[inline]
pub fn project_hopper<'a, T: Projectable>(
account: &'a AccountView<'a>,
expected_disc: u8,
) -> Result<Ref<'a, T>, ProgramError> {
project::<T>(account, 10, Some(expected_disc))
}
#[inline]
pub unsafe fn project_hopper_mut<'a, T: Projectable>(
account: &'a AccountView<'a>,
expected_disc: u8,
) -> Result<&'a mut T, ProgramError> {
unsafe { project_mut::<T>(account, 10, Some(expected_disc)) }
}
#[cfg(test)]
mod tests {
use super::*;
use crate::raw_account::RuntimeAccount;
use crate::NOT_BORROWED;
#[repr(C, align(8))]
struct Backing {
header: RuntimeAccount,
data: [u8; 16],
}
fn make_backing() -> Backing {
let header = RuntimeAccount {
borrow_state: NOT_BORROWED,
data_len: 16,
..RuntimeAccount::default()
};
Backing {
header,
data: [7u8; 16],
}
}
#[test]
fn projection_takes_a_shared_borrow_and_blocks_exclusive() {
let mut backing = make_backing();
let account = unsafe { AccountView::new_unchecked(&mut backing.header) };
{
let field = project::<u8>(&account, 0, None).unwrap();
assert_eq!(*field, 7);
assert!(account.try_borrow_mut().is_err());
}
assert!(account.try_borrow_mut().is_ok());
{
let _data = account.try_borrow_mut().unwrap();
assert!(project::<u8>(&account, 0, None).is_err());
}
assert!(project::<u8>(&account, 0, None).is_ok());
}
#[test]
fn project_bounds_and_disc_checks_run_before_borrowing() {
let mut backing = make_backing();
let account = unsafe { AccountView::new_unchecked(&mut backing.header) };
assert!(project::<[u8; 32]>(&account, 0, None).is_err());
assert!(project::<u8>(&account, 0, Some(9)).is_err());
assert!(account.try_borrow_mut().is_ok());
}
}