use std::cell::RefCell;
use std::panic::{self, AssertUnwindSafe};
use std::rc::Rc;
use super::odra_vm_state::OdraVmState;
use anyhow::Result;
use odra_core::callstack::CallstackElement;
use odra_core::casper_types::bytesrepr::{deserialize, deserialize_from_slice, serialize};
use odra_core::casper_types::system::auction::ValidatorBid;
use odra_core::casper_types::{CLType, CLValue, HashAddr, PackageHash, RuntimeArgs};
use odra_core::entry_point_callback::EntryPointsCaller;
use odra_core::prelude::*;
use odra_core::validator::ValidatorInfo;
use odra_core::CallDef;
use odra_core::EventError;
use odra_core::VmError;
use odra_core::{
callstack,
casper_types::{
bytesrepr::{Bytes, FromBytes, ToBytes},
PublicKey, SecretKey, U512
}
};
use odra_core::{ContractContainer, ContractRegister};
const NAMED_KEY_PREFIX: &str = "NAMED_KEY";
pub struct OdraVm {
state: Rc<RefCell<OdraVmState>>,
contract_register: Rc<RefCell<ContractRegister>>
}
impl Default for OdraVm {
fn default() -> Self {
Self {
state: Rc::new(RefCell::new(OdraVmState::default())),
contract_register: Rc::new(RefCell::new(ContractRegister::default()))
}
}
}
impl OdraVm {
pub fn new() -> Rc<Self> {
Rc::new(Self::default())
}
pub fn new_contract(
&self,
name: &str,
init_args: RuntimeArgs,
entry_points_caller: EntryPointsCaller
) -> Address {
let address = self.state.borrow_mut().next_contract_address();
{
let contract = ContractContainer::new(name, entry_points_caller);
let mut contract_register = self.contract_register.borrow_mut();
contract_register.add(address, contract);
self.state.borrow_mut().set_balance(address, U512::zero());
}
address
}
pub fn upgrade_contract(
&self,
name: &str,
contract_to_upgrade: Address,
upgrade_args: RuntimeArgs,
entry_points_caller: EntryPointsCaller
) -> Address {
let mut contract_register = self.contract_register.borrow_mut();
let contract = ContractContainer::new(name, entry_points_caller);
contract_register.add(contract_to_upgrade, contract);
contract_to_upgrade
}
pub(crate) fn post_install(&self, address: Address) {
self.contract_register.borrow_mut().post_install(&address);
}
pub fn call_contract(&self, address: Address, call_def: CallDef) -> Bytes {
let contract_name = self
.contract_register
.borrow()
.get(&address)
.map(|c| String::from(c.name()))
.unwrap_or(String::from("UnknownContractName"));
self.prepare_call(contract_name, address, &call_def);
if call_def.amount() > U512::zero() {
let status = self.checked_transfer_tokens(&self.caller(), &address, &call_def.amount());
if let Err(err) = status {
self.revert(err);
}
}
let result = self.contract_register.borrow().call(&address, call_def);
match result {
Err(err) => self.revert(err),
Ok(bytes) => self.handle_call_result(bytes)
}
}
pub fn revert(&self, error: OdraError) -> ! {
let mut revert_msg = String::from("");
if let CallstackElement::ContractCall {
contract_name,
address,
call_def
} = self.callstack_tip()
{
revert_msg = format!(
"{}({:?})::{}",
contract_name,
address,
call_def.entry_point()
);
}
let mut state = self.state.borrow_mut();
state.set_error(error.clone());
state.clear_callstack();
if state.is_in_caller_context() {
state.restore_snapshot();
}
drop(state);
panic!("Revert: {:?} - {}", error, revert_msg);
}
pub fn error(&self) -> Option<OdraError> {
self.state.borrow().error()
}
pub fn self_address(&self) -> Address {
self.state.borrow().callee()
}
pub fn caller(&self) -> Address {
self.state.borrow().caller()
}
pub fn read_stack_record(&self) -> String {
self.state.borrow().read_stack_record()
}
pub fn callee(&self) -> Address {
self.state.borrow().callee()
}
pub fn callstack_tip(&self) -> CallstackElement {
self.state.borrow().callstack_tip().clone()
}
pub fn get_named_arg(&self, name: &str) -> OdraResult<Vec<u8>> {
match self.state.borrow().callstack_tip() {
CallstackElement::Account(_) => todo!(),
CallstackElement::ContractCall { call_def, .. } => call_def
.args()
.get(name)
.map(|arg| arg.inner_bytes().to_vec())
.ok_or(OdraError::ExecutionError(ExecutionError::MissingArg))
}
}
pub fn set_caller(&self, caller: Address) {
self.state.borrow_mut().set_caller(caller);
}
pub fn set_var(&self, key: &[u8], value: Bytes) {
self.state.borrow_mut().set_var(key, value);
}
pub fn get_var(&self, key: &[u8]) -> Option<Bytes> {
let result = { self.state.borrow().get_var(key) };
match result {
Ok(result) => result,
Err(error) => {
self.state
.borrow_mut()
.set_error(Into::<ExecutionError>::into(error));
None
}
}
}
pub fn set_named_key(&self, name: &str, value: CLValue) {
let key = Self::key_of_named_key(name);
self.set_var(key.as_bytes(), Bytes::from(value.inner_bytes().as_slice()));
}
pub fn get_named_key(&self, name: &str) -> Option<Bytes> {
let key = Self::key_of_named_key(name);
self.get_var(key.as_bytes())
}
pub fn set_dict_value(&self, dict: &str, key: &[u8], value: CLValue) {
self.state.borrow_mut().set_dict_value(
dict.as_bytes(),
key,
Bytes::from(value.inner_bytes().as_slice())
);
}
pub fn remove_dictionary(&self, dictionary_name: &str) {
self.state
.borrow_mut()
.remove_dictionary(dictionary_name.as_bytes());
}
pub fn get_dict_value(&self, dict: &str, key: &[u8]) -> Option<Bytes> {
let result = { self.state.borrow().get_dict_value(dict.as_bytes(), key) };
match result {
Ok(result) => result,
Err(error) => {
self.state
.borrow_mut()
.set_error(Into::<ExecutionError>::into(error));
None
}
}
}
pub fn emit_event(&self, event_data: &Bytes) {
self.state.borrow_mut().emit_event(event_data);
}
pub fn emit_native_event(&self, event_data: &Bytes) {
self.state.borrow_mut().emit_native_event(event_data);
}
pub fn get_event(&self, address: &Address, index: u32) -> Result<Bytes, EventError> {
self.state.borrow().get_event(address, index)
}
pub fn get_native_event(&self, address: &Address, index: u32) -> Result<Bytes, EventError> {
self.state.borrow().get_native_event(address, index)
}
pub fn get_events_count(&self, address: &Address) -> Result<u32, EventError> {
self.state.borrow().get_events_count(address)
}
pub fn get_native_events_count(&self, address: &Address) -> Result<u32, EventError> {
self.state.borrow().get_native_events_count(address)
}
pub fn attach_value(&self, amount: U512) {
self.state.borrow_mut().attach_value(amount);
}
pub fn get_block_time(&self) -> u64 {
self.state.borrow().block_time()
}
pub fn advance_block_time_by(&self, milliseconds: u64) {
self.state.borrow_mut().advance_block_time_by(milliseconds)
}
pub fn advance_with_auctions(&self, milliseconds: u64) {
self.state.borrow_mut().advance_with_auctions(milliseconds)
}
pub fn attached_value(&self) -> U512 {
self.state.borrow().attached_value()
}
pub fn get_account(&self, n: usize) -> Address {
self.state.borrow().accounts.get(n).cloned().unwrap()
}
pub fn get_validator(&self, n: usize) -> PublicKey {
self.state
.borrow()
.validators
.iter()
.map(|a| a.0)
.nth(n)
.unwrap_or_else(|| panic!("Validator with index {} does not exist", n))
.clone()
}
pub fn balance_of(&self, address: &Address) -> U512 {
self.state.borrow().balance_of(address)
}
pub fn transfer_tokens(&self, to: &Address, amount: &U512) {
if amount.is_zero() {
return;
}
let from = &self.self_address();
let mut transfer_error = None;
{
let mut state = self.state.borrow_mut();
if state.transfer(from, to, amount).is_err() {
transfer_error = Some(OdraError::VmError(VmError::BalanceExceeded));
}
}
if let Some(result) = transfer_error {
self.revert(result);
}
}
pub fn checked_transfer_tokens(
&self,
from: &Address,
to: &Address,
amount: &U512
) -> OdraResult<()> {
if amount.is_zero() {
return Ok(());
}
let mut state = self.state.borrow_mut();
if state.transfer(from, to, amount).is_err() {
return Err(OdraError::VmError(VmError::BalanceExceeded));
}
Ok(())
}
pub fn self_balance(&self) -> U512 {
let address = self.self_address();
self.state.borrow().balance_of(&address)
}
pub fn public_key(&self, address: &Address) -> PublicKey {
self.state.borrow().public_key(address)
}
pub fn sign_message(&self, message: &Bytes, address: &Address) -> Bytes {
let public_key = self.public_key(address);
let signature = odra_core::casper_types::crypto::sign(
message,
self.state.borrow().secret_key(address),
&public_key
)
.to_bytes()
.unwrap();
signature.into()
}
pub fn delegated_amount(&self, delegator: Address, validator: PublicKey) -> U512 {
self.state.borrow().delegated_amount(validator, delegator)
}
pub fn get_validator_info(&self, validator: PublicKey) -> Option<ValidatorInfo> {
self.state.borrow().validators.get(&validator).cloned()
}
pub fn remove_validator(&self, index: usize) {
let validator = self.get_validator(index);
self.state.borrow_mut().remove_validator(validator);
}
pub fn delegate(&self, validator: PublicKey, delegator: Address, amount: U512) {
let mut state = self.state.borrow_mut();
state.delegate(validator, delegator, amount);
}
pub fn undelegate(&self, validator: PublicKey, delegator: Address, amount: U512) {
let mut state = self.state.borrow_mut();
state.undelegate(validator, delegator, amount);
}
pub fn auction_delay(&self) -> u64 {
self.state.borrow().auction_delay()
}
pub fn unbonding_delay(&self) -> u64 {
self.auction_delay() * 7
}
}
impl OdraVm {
fn prepare_call(&self, contract_name: String, address: Address, call_def: &CallDef) {
let mut state = self.state.borrow_mut();
if state.is_in_caller_context() {
state.take_snapshot();
state.clear_error();
}
let element = CallstackElement::new_contract_call(contract_name, address, call_def.clone());
state.push_callstack_element(element);
}
fn handle_call_result(&self, result: Bytes) -> Bytes {
let mut state = self.state.borrow_mut();
state.pop_callstack_element();
if state.is_in_caller_context() {
state.drop_snapshot();
}
result
}
fn key_of_named_key(name: &str) -> String {
let key = format!("{}_{}", NAMED_KEY_PREFIX, name);
key
}
}
#[cfg(test)]
mod tests {
use odra_core::callstack::CallstackElement;
use odra_core::casper_types::bytesrepr::{Bytes, ToBytes};
use odra_core::{
entry_point_callback::{EntryPoint, EntryPointsCaller},
utils::serialize
};
use std::collections::BTreeMap;
use odra_core::casper_types::bytesrepr::FromBytes;
use odra_core::casper_types::{CLValue, RuntimeArgs, U512};
use odra_core::host::HostEnv;
use odra_core::{prelude::*, CallDef, VmError};
use crate::vm::utils;
use crate::{OdraVm, OdraVmHost};
const TEST_ENTRY_POINT: &str = "abc";
#[test]
fn contracts_have_different_addresses() {
let instance = OdraVm::default();
let address1 =
instance.new_contract("A", RuntimeArgs::new(), test_caller(TEST_ENTRY_POINT));
let address2 =
instance.new_contract("B", RuntimeArgs::new(), test_caller(TEST_ENTRY_POINT));
assert_ne!(address1, address2);
}
#[test]
fn addresses_have_different_type() {
let instance = OdraVm::default();
let account_address = instance.get_account(0);
let contract_address = setup_contract(&instance, TEST_ENTRY_POINT);
assert!(contract_address.is_contract());
assert!(!account_address.is_contract());
}
#[test]
fn test_contract_call() {
let instance = OdraVm::default();
let contract_address = setup_contract(&instance, TEST_ENTRY_POINT);
let result = instance.call_contract(
contract_address,
CallDef::new(TEST_ENTRY_POINT, false, RuntimeArgs::new())
);
let expected_result: Bytes = test_call_result();
assert_eq!(result, expected_result);
}
#[test]
fn test_transfer() {
let instance = OdraVm::default();
let from = instance.get_account(0);
let from_balance = instance.balance_of(&from);
let to = instance.get_account(1);
let to_balance = instance.balance_of(&to);
let amount = U512::from(100);
instance.transfer_tokens(&to, &amount);
assert_eq!(instance.balance_of(&from), from_balance - amount);
assert_eq!(instance.balance_of(&to), to_balance + amount);
}
#[test]
#[should_panic]
fn test_transfer_too_much() {
let instance = OdraVm::default();
let from = instance.get_account(0);
let from_balance = instance.balance_of(&from);
let to = instance.get_account(1);
let to_balance = instance.balance_of(&to);
let amount = from_balance + 1;
instance.transfer_tokens(&to, &amount);
}
#[test]
fn test_call_non_existing_contract() {
let instance = OdraVm::default();
let address = utils::contract_address_from_u32(42);
let call_def = CallDef::new(TEST_ENTRY_POINT, false, RuntimeArgs::new());
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
instance.call_contract(address, call_def)
}));
assert_eq!(
instance.error(),
Some(OdraError::VmError(VmError::InvalidContractAddress))
);
}
#[test]
fn test_call_non_existing_entrypoint() {
let instance = OdraVm::default();
let invalid_entry_point_name = "aaa";
let contract_address = setup_contract(&instance, TEST_ENTRY_POINT);
let call_def = CallDef::new(invalid_entry_point_name, false, RuntimeArgs::new());
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
instance.call_contract(contract_address, call_def)
}));
assert_eq!(
instance.error(),
Some(OdraError::VmError(VmError::NoSuchMethod(
invalid_entry_point_name.to_string()
)))
);
}
#[test]
fn test_caller_switching() {
let instance = OdraVm::default();
let new_caller = utils::account_address_from_str("ff");
instance.set_caller(new_caller);
push_address(&instance, &new_caller);
assert_eq!(instance.caller(), new_caller);
}
#[test]
#[should_panic]
fn test_revert() {
let instance = OdraVm::default();
instance.revert(OdraError::user(1, "Test revert"));
}
#[test]
fn test_read_write_value() {
let instance = OdraVm::default();
let key = b"key";
let value = 32u8.to_bytes().map(Bytes::from).unwrap();
instance.set_var(key, value.clone());
assert_eq!(instance.get_var(key), Some(value));
assert_eq!(instance.get_var(b"other_key"), None);
}
#[test]
fn test_read_write_dict() {
let instance = OdraVm::default();
let dict = "dict";
let key = b"key";
let value = CLValue::from_t("value").unwrap();
instance.set_dict_value(dict, key, value.clone());
assert_eq!(
instance.get_dict_value(dict, key),
Some(Bytes::from(value.inner_bytes().as_slice()))
);
assert_eq!(instance.get_dict_value("other_dict", key), None);
assert_eq!(instance.get_dict_value(dict, b"other_key"), None);
}
#[test]
fn test_named_key() {
let instance = OdraVm::default();
let name = "name";
let value = CLValue::from_t("value").unwrap();
instance.set_named_key(name, value.clone());
assert_eq!(
instance.get_named_key(name),
Some(Bytes::from(value.inner_bytes().as_slice()))
);
assert_eq!(instance.get_named_key("other_name"), None);
}
#[test]
fn events() {
let instance = OdraVm::default();
let first_contract_address = utils::contract_address_from_u32(123);
push_address(&instance, &first_contract_address);
let first_event: Bytes = vec![1, 2, 3].into();
let second_event: Bytes = vec![4, 5, 6].into();
instance.emit_event(&first_event);
instance.emit_event(&second_event);
let second_contract_address = utils::contract_address_from_u32(321);
push_address(&instance, &second_contract_address);
let third_event: Bytes = vec![7, 8, 9].into();
let fourth_event: Bytes = vec![11, 22, 33].into();
instance.emit_event(&third_event);
instance.emit_event(&fourth_event);
assert_eq!(
instance.get_event(&first_contract_address, 0),
Ok(first_event)
);
assert_eq!(
instance.get_event(&first_contract_address, 1),
Ok(second_event)
);
assert_eq!(
instance.get_event(&second_contract_address, 0),
Ok(third_event)
);
assert_eq!(
instance.get_event(&second_contract_address, 1),
Ok(fourth_event)
);
}
#[test]
fn test_current_contract_address() {
let instance = OdraVm::default();
let contract_address = setup_contract(&instance, TEST_ENTRY_POINT);
let contract_address = utils::contract_address_from_u32(100);
push_address(&instance, &contract_address);
assert_eq!(instance.self_address(), contract_address);
}
#[test]
fn test_call_contract_with_amount() {
let instance = OdraVm::default();
let contract_address = setup_contract(&instance, TEST_ENTRY_POINT);
let caller = instance.get_account(0);
let caller_balance = instance.balance_of(&caller);
let call_def =
CallDef::new(TEST_ENTRY_POINT, false, RuntimeArgs::new()).with_amount(caller_balance);
instance.call_contract(contract_address, call_def);
assert_eq!(instance.balance_of(&contract_address), caller_balance);
assert_eq!(instance.balance_of(&caller), U512::zero());
}
#[test]
#[should_panic(expected = "VmError(BalanceExceeded)")]
fn test_call_contract_with_amount_exceeding_balance() {
let instance = OdraVm::default();
let contract_address = setup_contract(&instance, TEST_ENTRY_POINT);
let caller = instance.get_account(0);
let caller_balance = instance.balance_of(&caller);
let call_def = CallDef::new(TEST_ENTRY_POINT, false, RuntimeArgs::new())
.with_amount(caller_balance + 1);
instance.call_contract(contract_address, call_def);
}
fn push_address(vm: &OdraVm, address: &Address) {
let element = CallstackElement::new_account(*address);
vm.state.borrow_mut().push_callstack_element(element);
}
fn test_call_result() -> Bytes {
vec![1, 1, 0, 0].into()
}
fn setup_contract(instance: &OdraVm, entry_point_name: &str) -> Address {
let caller = test_caller(entry_point_name);
instance.new_contract("contract", RuntimeArgs::new(), caller)
}
fn test_caller(entry_point_name: &str) -> EntryPointsCaller {
let vm = OdraVm::new();
let host_env = OdraVmHost::new(vm);
let env = HostEnv::new(host_env);
let entry_point = EntryPoint::new_payable(String::from(entry_point_name), vec![]);
EntryPointsCaller::new(env, vec![entry_point], |_, _| Ok(test_call_result()))
}
}