magicsvm 0.2.1

A fast and lightweight Solana + MagicBlock VM simulator for testing solana programs
use {
    crate::magic::EPHEMERAL_VAULT_ID,
    solana_account::{ReadableAccount, WritableAccount},
    solana_program_runtime::invoke_context::InvokeContext,
    solana_sdk_ids::system_program,
    solana_transaction::InstructionError,
    solana_transaction_context::{transaction::TransactionContext, IndexOfAccount},
};

const EPHEMERAL_SPONSOR_IDX: IndexOfAccount = 0;
const EPHEMERAL_ACCOUNT_IDX: IndexOfAccount = 1;
const EPHEMERAL_VAULT_IDX: IndexOfAccount = 2;
const MAX_EPHEMERAL_DATA_LEN: u32 = 10 * 1024 * 1024;

pub(super) fn process_create_ephemeral_account(
    invoke_context: &InvokeContext,
    data_len: u32,
) -> Result<(), InstructionError> {
    if data_len > MAX_EPHEMERAL_DATA_LEN {
        return Err(InstructionError::InvalidArgument);
    }

    let transaction_context = &invoke_context.transaction_context;
    let caller_program_id = validate_ephemeral_common(transaction_context)?;
    validate_instruction_account_signer(transaction_context, EPHEMERAL_ACCOUNT_IDX)?;

    let ephemeral_index = transaction_context
        .get_current_instruction_context()?
        .get_index_of_instruction_account_in_transaction(EPHEMERAL_ACCOUNT_IDX)?;
    let mut account = transaction_context
        .accounts()
        .try_borrow_mut(ephemeral_index)?;
    if account.lamports() != 0 || account.owner() != &system_program::ID {
        return Err(InstructionError::InvalidAccountData);
    }
    account.set_owner(caller_program_id);
    account.set_data_from_slice(&vec![0u8; data_len as usize]);
    Ok(())
}

pub(super) fn process_resize_ephemeral_account(
    invoke_context: &InvokeContext,
    new_data_len: u32,
) -> Result<(), InstructionError> {
    if new_data_len > MAX_EPHEMERAL_DATA_LEN {
        return Err(InstructionError::InvalidArgument);
    }

    let transaction_context = &invoke_context.transaction_context;
    let caller_program_id = validate_ephemeral_common(transaction_context)?;
    validate_existing_ephemeral(transaction_context, caller_program_id)?;

    let ephemeral_index = transaction_context
        .get_current_instruction_context()?
        .get_index_of_instruction_account_in_transaction(EPHEMERAL_ACCOUNT_IDX)?;
    let mut account = transaction_context
        .accounts()
        .try_borrow_mut(ephemeral_index)?;
    let mut data = account.data().to_vec();
    data.resize(new_data_len as usize, 0);
    account.set_data_from_slice(&data);
    Ok(())
}

pub(super) fn process_close_ephemeral_account(
    invoke_context: &InvokeContext,
) -> Result<(), InstructionError> {
    let transaction_context = &invoke_context.transaction_context;
    let caller_program_id = validate_ephemeral_common(transaction_context)?;
    validate_existing_ephemeral(transaction_context, caller_program_id)?;
    Ok(())
}

fn validate_ephemeral_common(
    transaction_context: &TransactionContext<'_>,
) -> Result<solana_address::Address, InstructionError> {
    let caller_program_id =
        get_caller_program_id(transaction_context).ok_or(InstructionError::IncorrectProgramId)?;
    validate_instruction_account_signer(transaction_context, EPHEMERAL_SPONSOR_IDX)?;
    if instruction_account_key(transaction_context, EPHEMERAL_VAULT_IDX)? != &EPHEMERAL_VAULT_ID {
        return Err(InstructionError::InvalidAccountData);
    }
    Ok(caller_program_id)
}

fn validate_existing_ephemeral(
    transaction_context: &TransactionContext<'_>,
    caller_program_id: solana_address::Address,
) -> Result<u32, InstructionError> {
    let instruction_context = transaction_context.get_current_instruction_context()?;
    let ephemeral = instruction_context.try_borrow_instruction_account(EPHEMERAL_ACCOUNT_IDX)?;
    if ephemeral.get_owner() != &caller_program_id {
        return Err(InstructionError::InvalidAccountOwner);
    }
    ephemeral
        .get_data()
        .len()
        .try_into()
        .map_err(|_| InstructionError::ArithmeticOverflow)
}

fn validate_instruction_account_signer(
    transaction_context: &TransactionContext<'_>,
    account_index: IndexOfAccount,
) -> Result<(), InstructionError> {
    if !transaction_context
        .get_current_instruction_context()?
        .is_instruction_account_signer(account_index)?
    {
        return Err(InstructionError::MissingRequiredSignature);
    }
    Ok(())
}

fn get_caller_program_id(
    transaction_context: &TransactionContext<'_>,
) -> Option<solana_address::Address> {
    let current_instruction_context = transaction_context.get_current_instruction_context().ok()?;
    let current_instruction_index = current_instruction_context.get_index_in_trace();
    let current_stack_height = current_instruction_context.get_stack_height();
    if current_stack_height <= 1 {
        return None;
    }

    (0..current_instruction_index).rev().find_map(|index| {
        let instruction_context = transaction_context
            .get_instruction_context_at_index_in_trace(index)
            .ok()?;
        (instruction_context.get_stack_height() < current_stack_height)
            .then(|| instruction_context.get_program_key().ok().copied())
            .flatten()
    })
}

fn instruction_account_key<'a, 'ix_data>(
    transaction_context: &'a TransactionContext<'ix_data>,
    account_index: IndexOfAccount,
) -> Result<&'a solana_address::Address, InstructionError> {
    let transaction_index = transaction_context
        .get_current_instruction_context()?
        .get_index_of_instruction_account_in_transaction(account_index)?;
    transaction_context.get_key_of_account_at_index(transaction_index)
}