use core::marker::PhantomData;
use crate::account::AccountView;
use crate::error::ProgramError;
use crate::layout::LayoutContract;
use crate::ProgramResult;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct BehaviorWrite {
pub offset: u32,
pub size: u32,
}
impl BehaviorWrite {
#[inline(always)]
pub const fn new(offset: u32, size: u32) -> Self {
Self { offset, size }
}
}
pub trait HopperBehavior<T: LayoutContract> {
type Args;
type CheckOutput;
const RUN_CHECK: bool = true;
const RUN_UPDATE: bool = false;
const RUN_EXIT: bool = false;
const REQUIRES_MUT: bool = Self::RUN_UPDATE || Self::RUN_EXIT;
const WRITES: &'static [BehaviorWrite] = &[];
fn check(
view: &AccountView<'_>,
state: &T,
args: &Self::Args,
) -> Result<Self::CheckOutput, ProgramError> {
let _ = (view, state, args);
Err(ProgramError::InvalidArgument)
}
fn update(view: &AccountView<'_>, state: &mut T, args: &Self::Args) -> ProgramResult {
let _ = (view, state, args);
Ok(())
}
fn exit(view: &AccountView<'_>, args: &Self::Args) -> ProgramResult {
let _ = (view, args);
Ok(())
}
}
pub struct BehaviorChecked<B, O> {
pub output: O,
_behavior: PhantomData<B>,
}
impl<B, O> BehaviorChecked<B, O> {
#[inline(always)]
fn new(output: O) -> Self {
Self {
output,
_behavior: PhantomData,
}
}
}
#[inline]
pub fn run_check<B, T>(
view: &AccountView<'_>,
args: &B::Args,
) -> Result<BehaviorChecked<B, B::CheckOutput>, ProgramError>
where
T: LayoutContract + crate::Pod,
B: HopperBehavior<T>,
{
if !B::RUN_CHECK {
return Err(ProgramError::InvalidArgument);
}
let state = view.load::<T>()?;
let output = B::check(view, &state, args)?;
Ok(BehaviorChecked::new(output))
}
#[inline]
pub fn run_update<B, T>(
view: &AccountView<'_>,
args: &B::Args,
_proof: &BehaviorChecked<B, B::CheckOutput>,
) -> ProgramResult
where
T: LayoutContract + crate::Pod,
B: HopperBehavior<T>,
{
if !B::RUN_UPDATE {
return Ok(());
}
let mut state = view.load_mut::<T>()?;
B::update(view, &mut state, args)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::layout::HopperHeader;
use hopper_native::{
AccountView as NativeAccountView, Address as NativeAddress, RuntimeAccount, NOT_BORROWED,
};
#[repr(C)]
#[derive(Clone, Copy)]
struct FeeVault {
collected_bps: [u8; 2],
_pad: [u8; 6],
}
unsafe impl crate::Zeroable for FeeVault {}
unsafe impl crate::Pod for FeeVault {}
impl crate::field_map::FieldMap for FeeVault {
const FIELDS: &'static [crate::field_map::FieldInfo] = &[crate::field_map::FieldInfo::new(
"collected_bps",
HopperHeader::SIZE,
2,
)];
}
impl LayoutContract for FeeVault {
const DISC: u8 = 42;
const VERSION: u8 = 1;
const LAYOUT_ID: [u8; 8] = [0x42; 8];
const SIZE: usize = HopperHeader::SIZE + core::mem::size_of::<Self>();
}
struct FeeCap;
struct FeeCapArgs {
max_bps: u16,
}
impl HopperBehavior<FeeVault> for FeeCap {
type Args = FeeCapArgs;
type CheckOutput = u16;
const WRITES: &'static [BehaviorWrite] = &[BehaviorWrite::new(
HopperHeader::SIZE as u32,
2, )];
const RUN_UPDATE: bool = true;
fn check(
_view: &AccountView<'_>,
state: &FeeVault,
args: &Self::Args,
) -> Result<u16, ProgramError> {
let bps = u16::from_le_bytes(state.collected_bps);
if bps > args.max_bps {
return Err(ProgramError::InvalidAccountData);
}
Ok(bps)
}
fn update(
_view: &AccountView<'_>,
state: &mut FeeVault,
args: &Self::Args,
) -> ProgramResult {
let bps = u16::from_le_bytes(state.collected_bps).min(args.max_bps);
state.collected_bps = bps.to_le_bytes();
Ok(())
}
}
fn make_vault(bps: u16) -> (std::vec::Vec<u64>, AccountView<'static>) {
let data_len = FeeVault::SIZE;
let mut backing = std::vec![0u64; (RuntimeAccount::SIZE + data_len).div_ceil(8)];
let raw = backing.as_mut_ptr() as *mut RuntimeAccount;
unsafe {
raw.write(RuntimeAccount {
borrow_state: NOT_BORROWED,
is_signer: 0,
is_writable: 1,
executable: 0,
resize_delta: 0,
address: NativeAddress::new_from_array([1; 32]),
owner: NativeAddress::new_from_array([2; 32]),
lamports: 1,
data_len: data_len as u64,
});
}
let backend = unsafe { NativeAccountView::new_unchecked(raw) };
let view = AccountView::from_backend(backend);
{
let mut d = view.try_borrow_mut().unwrap();
crate::layout::init_header::<FeeVault>(&mut d).unwrap();
d[HopperHeader::SIZE..HopperHeader::SIZE + 2].copy_from_slice(&bps.to_le_bytes());
}
(backing, view)
}
#[test]
fn check_mints_proof_with_payload_and_rejects_violations() {
let (_b, vault) = make_vault(25);
let args = FeeCapArgs { max_bps: 30 };
let proof = run_check::<FeeCap, FeeVault>(&vault, &args).unwrap();
assert_eq!(proof.output, 25);
let (_b2, hot) = make_vault(31);
assert!(run_check::<FeeCap, FeeVault>(&hot, &args).is_err());
}
#[test]
fn update_requires_the_proof_and_applies_declared_writes() {
let (_b, vault) = make_vault(30);
let args = FeeCapArgs { max_bps: 30 };
let proof = run_check::<FeeCap, FeeVault>(&vault, &args).unwrap();
run_update::<FeeCap, FeeVault>(&vault, &args, &proof).unwrap();
let state = vault.load::<FeeVault>().unwrap();
assert_eq!(u16::from_le_bytes(state.collected_bps), 30);
}
#[test]
fn write_contribution_is_declared_for_strict_writes_folding() {
assert_eq!(
<FeeCap as HopperBehavior<FeeVault>>::WRITES,
&[BehaviorWrite::new(HopperHeader::SIZE as u32, 2)]
);
const { assert!(<FeeCap as HopperBehavior<FeeVault>>::REQUIRES_MUT) };
}
}