use crate::address::Address;
use crate::error::ProgramError;
use crate::project::Projectable;
use core::mem::MaybeUninit;
#[cfg(feature = "cpi")]
use crate::instruction::{InstructionView, Signer};
pub const MAX_RETURN_DATA: usize = 1024;
pub struct ReturnData {
buf: [MaybeUninit<u8>; MAX_RETURN_DATA],
len: usize,
program_id: Address,
}
impl ReturnData {
#[inline(always)]
pub fn data(&self) -> &[u8] {
debug_assert!(self.len <= MAX_RETURN_DATA);
unsafe { core::slice::from_raw_parts(self.buf.as_ptr() as *const u8, self.len) }
}
#[inline(always)]
pub fn program_id(&self) -> &Address {
&self.program_id
}
#[inline(always)]
pub fn len(&self) -> usize {
self.len
}
#[inline(always)]
pub fn is_empty(&self) -> bool {
self.len == 0
}
#[inline]
pub fn as_type<T: Projectable>(&self) -> Result<&T, ProgramError> {
let size = core::mem::size_of::<T>();
if self.len < size {
return Err(ProgramError::AccountDataTooSmall);
}
let data = self.data();
let align = core::mem::align_of::<T>();
let ptr = data.as_ptr();
if !(ptr as usize).is_multiple_of(align) {
return Err(ProgramError::InvalidAccountData);
}
Ok(unsafe { &*(ptr as *const T) })
}
#[inline]
pub fn as_type_from<T: Projectable>(
&self,
expected_program: &Address,
) -> Result<&T, ProgramError> {
if !crate::address::address_eq(self.program_id(), expected_program) {
return Err(ProgramError::IncorrectProgramId);
}
self.as_type::<T>()
}
#[inline]
pub fn as_u64(&self) -> Result<u64, ProgramError> {
if self.len < 8 {
return Err(ProgramError::AccountDataTooSmall);
}
let mut bytes = [0u8; 8];
bytes.copy_from_slice(&self.data()[..8]);
Ok(u64::from_le_bytes(bytes))
}
#[inline]
pub fn as_u32(&self) -> Result<u32, ProgramError> {
if self.len < 4 {
return Err(ProgramError::AccountDataTooSmall);
}
let mut bytes = [0u8; 4];
bytes.copy_from_slice(&self.data()[..4]);
Ok(u32::from_le_bytes(bytes))
}
}
#[inline]
pub fn get_return_data() -> Option<ReturnData> {
#[allow(unused_mut)]
let mut rd = ReturnData {
buf: [const { MaybeUninit::uninit() }; MAX_RETURN_DATA],
len: 0,
program_id: Address::default(),
};
#[cfg(target_os = "solana")]
{
let actual_len = unsafe {
crate::syscalls::sol_get_return_data(
rd.buf.as_mut_ptr() as *mut u8,
MAX_RETURN_DATA as u64,
rd.program_id.0.as_mut_ptr(),
)
};
rd.len = (actual_len as usize).min(MAX_RETURN_DATA);
}
#[cfg(not(target_os = "solana"))]
{
}
if rd.len == 0 {
None
} else {
Some(rd)
}
}
#[cfg(feature = "cpi")]
#[inline]
pub fn invoke_and_read<T: Projectable, const ACCOUNTS: usize>(
instruction: &InstructionView<'_, '_, '_, '_>,
account_views: &[&crate::account_view::AccountView<'_>; ACCOUNTS],
signers_seeds: &[Signer<'_, '_>],
) -> Result<ReturnData, ProgramError> {
crate::cpi::invoke_signed::<ACCOUNTS>(instruction, account_views, signers_seeds)?;
let returned = get_return_data().ok_or(ProgramError::InvalidAccountData)?;
returned.as_type_from::<T>(instruction.program_id)?;
Ok(returned)
}
#[cfg(test)]
impl ReturnData {
fn test_snapshot(bytes: &[u8], program_id: Address) -> Self {
assert!(bytes.len() <= MAX_RETURN_DATA);
let mut buf = [const { MaybeUninit::uninit() }; MAX_RETURN_DATA];
for (dst, src) in buf.iter_mut().zip(bytes) {
dst.write(*src);
}
ReturnData {
buf,
len: bytes.len(),
program_id,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn typed_return_requires_the_expected_producer_and_initialized_type() {
let expected = Address::new_from_array([1; 32]);
let nested = Address::new_from_array([2; 32]);
let correct = ReturnData::test_snapshot(&42u64.to_le_bytes(), expected.clone());
assert_eq!(*correct.as_type_from::<u64>(&expected).unwrap(), 42);
assert_eq!(
correct.as_type_from::<u64>(&nested),
Err(ProgramError::IncorrectProgramId)
);
let short = ReturnData::test_snapshot(&[42], expected.clone());
assert_eq!(
short.as_type_from::<u64>(&expected),
Err(ProgramError::AccountDataTooSmall)
);
let forwarded = ReturnData::test_snapshot(&42u64.to_le_bytes(), nested);
assert_eq!(
forwarded.as_type_from::<u64>(&expected),
Err(ProgramError::IncorrectProgramId)
);
}
#[test]
fn offchain_get_return_data_is_none() {
assert!(get_return_data().is_none());
}
#[test]
fn data_exposes_exactly_the_written_prefix() {
let payload = [0xAB, 0xCD, 0xEF];
let rd = ReturnData::test_snapshot(&payload, Address::default());
assert_eq!(rd.data(), &payload);
assert_eq!(rd.len(), payload.len());
assert!(!rd.is_empty());
}
#[test]
fn as_u64_and_as_u32_never_read_past_the_prefix() {
let short = ReturnData::test_snapshot(&[1, 2, 3], Address::default());
assert!(short.as_u64().is_err());
assert!(short.as_u32().is_err());
let rd = ReturnData::test_snapshot(&7u64.to_le_bytes(), Address::default());
assert_eq!(rd.as_u64().unwrap(), 7);
assert_eq!(rd.as_u32().unwrap(), 7);
}
#[test]
fn as_type_length_checks_against_the_prefix() {
let rd = ReturnData::test_snapshot(&[5u8], Address::default());
assert!(rd.as_type::<u64>().is_err());
assert_eq!(*rd.as_type::<u8>().unwrap(), 5);
}
}