diff --git a/README.md b/README.md index 00e0f90..9dd63f5 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,7 @@ -# ERC4626 inflation attack free +# Cairo ERC4626 Component [inflation attack free] -An implementation of ERC4626 in Cairo. +An implementation of ERC4626 Component in Cairo. Forked from https://github.com/0xEniotna/ERC4626. Below is the original readme. +-------- I used the [OZ solidity](https://github.com/OpenZeppelin/openzeppelin-contracts/blob/master/contracts/token/ERC20/extensions/ERC4626.sol#L239) implementation. It is itself inspired from YieldBox codebase that has an inflation attack protection. Shares are virtually minted which reduces the issue. diff --git a/Scarb.toml b/Scarb.toml index 2d94f75..97a7380 100644 --- a/Scarb.toml +++ b/Scarb.toml @@ -1,11 +1,16 @@ [package] name = "erc4626" version = "0.1.0" +edition = "2023_11" [dependencies] -openzeppelin = { git = "https://github.com/OpenZeppelin/cairo-contracts.git"} -snforge_std = { git = "https://github.com/foundry-rs/starknet-foundry", tag = "v0.14.0" } -starknet = "2.4.3" +openzeppelin = "0.20.0" +starknet = ">=2.6.0" + +[dev-dependencies] +snforge_std = { git = "https://github.com/foundry-rs/starknet-foundry", tag = "v0.32.0" } + +[lib] [[target.starknet-contract]] sierra = true diff --git a/src/ERC4626.cairo b/src/ERC4626.cairo deleted file mode 100644 index 864240f..0000000 --- a/src/ERC4626.cairo +++ /dev/null @@ -1,5 +0,0 @@ -mod erc4626; -mod interface; -use erc4626::ERC4626; - -use interface::{IERC4626, IERC4626Dispatcher, IERC4626DispatcherTrait}; diff --git a/src/erc4626.cairo b/src/erc4626.cairo new file mode 100644 index 0000000..151f5cd --- /dev/null +++ b/src/erc4626.cairo @@ -0,0 +1,2 @@ +mod erc4626; +mod interface; diff --git a/src/erc4626/erc4626.cairo b/src/erc4626/erc4626.cairo index 91bf1cc..ab0d04d 100644 --- a/src/erc4626/erc4626.cairo +++ b/src/erc4626/erc4626.cairo @@ -1,37 +1,40 @@ -#[starknet::contract] -mod ERC4626 { +use starknet::ContractAddress; + +#[starknet::component] +pub mod ERC4626Component { + use openzeppelin::introspection::interface::{ISRC5Dispatcher, ISRC5DispatcherTrait}; + use openzeppelin::introspection::src5::SRC5Component::InternalTrait as SRC5InternalTrait; + use openzeppelin::introspection::src5::SRC5Component; + use erc4626::erc4626::interface::{ - IERC4626, IERC4626Additional, IERC4626Snake, IERC4626Camel, IERC4626Metadata + IERC4626Additional, IERC4626Snake, IERC4626Camel, IERC4626Metadata }; + use core::num::traits::Bounded; use erc4626::utils::{pow_256}; - use integer::BoundedU256; use openzeppelin::token::erc20::interface::{ - IERC20, IERC20Metadata, ERC20ABIDispatcher, ERC20ABIDispatcherTrait + IERC20, IERC20Metadata, ERC20ABIDispatcher, ERC20ABIDispatcherTrait, }; - use openzeppelin::token::erc20::{ERC20Component, ERC20Component::Errors as ERC20Errors}; + use openzeppelin::token::erc20::{ + ERC20Component, + ERC20HooksEmptyImpl, + ERC20Component::Errors as ERC20Errors + }; + use openzeppelin::token::erc20::ERC20Component::InternalTrait as ERC20InternalTrait; use starknet::{ContractAddress, get_caller_address, get_contract_address}; - component!(path: ERC20Component, storage: erc20, event: ERC20Event); - impl ERC20InternalImpl = ERC20Component::InternalImpl; - impl ERC20MetadataImpl = ERC20Component::ERC20MetadataImpl; - #[storage] struct Storage { asset: ContractAddress, underlying_decimals: u8, offset: u8, - #[substorage(v0)] - erc20: ERC20Component::Storage, } #[event] #[derive(Drop, starknet::Event)] - enum Event { + pub enum Event { Deposit: Deposit, Withdraw: Withdraw, - #[flat] - ERC20Event: ERC20Component::Event, } #[derive(Drop, starknet::Event)] @@ -57,40 +60,66 @@ mod ERC4626 { } mod Errors { - const EXCEEDED_MAX_DEPOSIT: felt252 = 'ERC4626: exceeded max deposit'; - const EXCEEDED_MAX_MINT: felt252 = 'ERC4626: exceeded max mint'; - const EXCEEDED_MAX_REDEEM: felt252 = 'ERC4626: exceeded max redeem'; - const EXCEEDED_MAX_WITHDRAW: felt252 = 'ERC4626: exceeded max withdraw'; + pub const EXCEEDED_MAX_DEPOSIT: felt252 = 'ERC4626: exceeded max deposit'; + pub const EXCEEDED_MAX_MINT: felt252 = 'ERC4626: exceeded max mint'; + pub const EXCEEDED_MAX_REDEEM: felt252 = 'ERC4626: exceeded max redeem'; + pub const EXCEEDED_MAX_WITHDRAW: felt252 = 'ERC4626: exceeded max withdraw'; } - #[constructor] - fn constructor( - ref self: ContractState, asset: ContractAddress, name: felt252, symbol: felt252, offset: u8 - ) { - let dispatcher = ERC20ABIDispatcher { contract_address: asset }; - self.offset.write(offset); - let decimals = dispatcher.decimals(); - self.erc20.initializer(name, symbol); - self.asset.write(asset); - self.underlying_decimals.write(decimals); + pub trait ERC4626HooksTrait { + fn before_deposit( + ref self: ComponentState, + caller: ContractAddress, + receiver: ContractAddress, + assets: u256, + shares: u256, + ); + fn after_deposit( + ref self: ComponentState, + caller: ContractAddress, + receiver: ContractAddress, + assets: u256, + shares: u256, + ); + fn before_withdraw( + ref self: ComponentState, + caller: ContractAddress, + receiver: ContractAddress, + owner: ContractAddress, + assets: u256, + shares: u256 + ); + fn after_withdraw( + ref self: ComponentState, + caller: ContractAddress, + receiver: ContractAddress, + owner: ContractAddress, + assets: u256, + shares: u256 + ); } - - #[abi(embed_v0)] - impl ERC4626Additional of IERC4626Additional { - fn asset(self: @ContractState) -> ContractAddress { + #[embeddable_as(ERC4626AdditionalImpl)] + impl ERC4626Additional< + TContractState, +HasComponent, + +ERC20Component::HasComponent, + +SRC5Component::HasComponent, + +ERC4626HooksTrait, + +Drop + > of IERC4626Additional> { + fn asset(self: @ComponentState) -> ContractAddress { self.asset.read() } - fn convert_to_assets(self: @ContractState, shares: u256) -> u256 { + fn convert_to_assets(self: @ComponentState, shares: u256) -> u256 { self._convert_to_assets(shares, false) } - fn convert_to_shares(self: @ContractState, assets: u256) -> u256 { + fn convert_to_shares(self: @ComponentState, assets: u256) -> u256 { self._convert_to_shares(assets, false) } - fn deposit(ref self: ContractState, assets: u256, receiver: ContractAddress) -> u256 { + fn deposit(ref self: ComponentState, assets: u256, receiver: ContractAddress) -> u256 { let max_assets = self.max_deposit(receiver); assert(max_assets >= assets, Errors::EXCEEDED_MAX_DEPOSIT); @@ -101,24 +130,24 @@ mod ERC4626 { shares } - fn max_deposit(self: @ContractState, address: ContractAddress) -> u256 { - BoundedU256::max() + fn max_deposit(self: @ComponentState, address: ContractAddress) -> u256 { + Bounded::::MAX } - fn max_mint(self: @ContractState, receiver: ContractAddress) -> u256 { - BoundedU256::max() + fn max_mint(self: @ComponentState, receiver: ContractAddress) -> u256 { + Bounded::::MAX } - fn max_redeem(self: @ContractState, owner: ContractAddress) -> u256 { + fn max_redeem(self: @ComponentState, owner: ContractAddress) -> u256 { self.balance_of(owner) } - fn max_withdraw(self: @ContractState, owner: ContractAddress) -> u256 { + fn max_withdraw(self: @ComponentState, owner: ContractAddress) -> u256 { let balance = self.balance_of(owner); self._convert_to_assets(balance, false) } - fn mint(ref self: ContractState, shares: u256, receiver: ContractAddress) -> u256 { + fn mint(ref self: ComponentState, shares: u256, receiver: ContractAddress) -> u256 { let max_shares = self.max_mint(receiver); assert(max_shares >= shares, Errors::EXCEEDED_MAX_MINT); @@ -129,24 +158,24 @@ mod ERC4626 { assets } - fn preview_deposit(self: @ContractState, assets: u256) -> u256 { + fn preview_deposit(self: @ComponentState, assets: u256) -> u256 { self._convert_to_shares(assets, false) } - fn preview_mint(self: @ContractState, shares: u256) -> u256 { + fn preview_mint(self: @ComponentState, shares: u256) -> u256 { self._convert_to_assets(shares, true) } - fn preview_redeem(self: @ContractState, shares: u256) -> u256 { + fn preview_redeem(self: @ComponentState, shares: u256) -> u256 { self._convert_to_assets(shares, false) } - fn preview_withdraw(self: @ContractState, assets: u256) -> u256 { + fn preview_withdraw(self: @ComponentState, assets: u256) -> u256 { self._convert_to_shares(assets, true) } fn redeem( - ref self: ContractState, shares: u256, receiver: ContractAddress, owner: ContractAddress + ref self: ComponentState, shares: u256, receiver: ContractAddress, owner: ContractAddress ) -> u256 { let max_shares = self.max_redeem(owner); assert(shares <= max_shares, Errors::EXCEEDED_MAX_REDEEM); @@ -157,13 +186,13 @@ mod ERC4626 { assets } - fn total_assets(self: @ContractState) -> u256 { + fn total_assets(self: @ComponentState) -> u256 { let dispatcher = ERC20ABIDispatcher { contract_address: self.asset.read() }; - dispatcher.balance_of(get_contract_address()) + dispatcher.balanceOf(get_contract_address()) } fn withdraw( - ref self: ContractState, assets: u256, receiver: ContractAddress, owner: ContractAddress + ref self: ComponentState, assets: u256, receiver: ContractAddress, owner: ContractAddress ) -> u256 { let max_assets = self.max_withdraw(owner); assert(assets <= max_assets, Errors::EXCEEDED_MAX_WITHDRAW); @@ -177,64 +206,90 @@ mod ERC4626 { } - #[abi(embed_v0)] - impl MetadataEntrypoints of IERC4626Metadata { - fn name(self: @ContractState) -> felt252 { - self.erc20.name() + #[embeddable_as(MetadataEntrypointsImpl)] + impl MetadataEntrypoints< + TContractState, +HasComponent, + impl erc20: ERC20Component::HasComponent, + +SRC5Component::HasComponent, + +ERC4626HooksTrait, + +Drop + > of IERC4626Metadata> { + fn name(self: @ComponentState) -> ByteArray { + let erc20_comp = get_dep_component!(ref self, erc20); + erc20_comp.name() } - fn symbol(self: @ContractState) -> felt252 { - self.erc20.symbol() + fn symbol(self: @ComponentState) -> ByteArray { + let erc20_comp = get_dep_component!(ref self, erc20); + erc20_comp.symbol() } - fn decimals(self: @ContractState) -> u8 { + fn decimals(self: @ComponentState) -> u8 { self.underlying_decimals.read() + self._decimals_offset() } } - #[abi(embed_v0)] - impl SnakeEntrypoints of IERC4626Snake { - fn total_supply(self: @ContractState) -> u256 { - self.erc20.total_supply() + #[embeddable_as(SnakeEntrypointsImpl)] + impl SnakeEntrypoints< + TContractState, +HasComponent, + impl erc20: ERC20Component::HasComponent, + +SRC5Component::HasComponent, + +ERC4626HooksTrait, + +Drop + > of IERC4626Snake> { + fn total_supply(self: @ComponentState) -> u256 { + let erc20_comp = get_dep_component!(ref self, erc20); + erc20_comp.total_supply() } - fn balance_of(self: @ContractState, account: ContractAddress) -> u256 { - self.erc20.balance_of(account) + fn balance_of(self: @ComponentState, account: ContractAddress) -> u256 { + let erc20_comp = get_dep_component!(ref self, erc20); + erc20_comp.balance_of(account) } fn allowance( - self: @ContractState, owner: ContractAddress, spender: ContractAddress + self: @ComponentState, owner: ContractAddress, spender: ContractAddress ) -> u256 { - self.erc20.allowance(owner, spender) + let erc20_comp = get_dep_component!(ref self, erc20); + erc20_comp.allowance(owner, spender) } - fn transfer(ref self: ContractState, recipient: ContractAddress, amount: u256) -> bool { - self.erc20.transfer(recipient, amount) + fn transfer(ref self: ComponentState, recipient: ContractAddress, amount: u256) -> bool { + let mut erc20_comp_mut = get_dep_component_mut!(ref self, erc20); + erc20_comp_mut.transfer(recipient, amount) } fn transfer_from( - ref self: ContractState, + ref self: ComponentState, sender: ContractAddress, recipient: ContractAddress, amount: u256 ) -> bool { - self.erc20.transfer_from(sender, recipient, amount) + let mut erc20_comp_mut = get_dep_component_mut!(ref self, erc20); + erc20_comp_mut.transfer_from(sender, recipient, amount) } - fn approve(ref self: ContractState, spender: ContractAddress, amount: u256) -> bool { - self.erc20.approve(spender, amount) + fn approve(ref self: ComponentState, spender: ContractAddress, amount: u256) -> bool { + let mut erc20_comp_mut = get_dep_component_mut!(ref self, erc20); + erc20_comp_mut.approve(spender, amount) } } - #[abi(embed_v0)] - impl CamelEntrypoints of IERC4626Camel { - fn totalSupply(self: @ContractState) -> u256 { + #[embeddable_as(CamelEntrypointsImpl)] + impl CamelEntrypoints< + TContractState, +HasComponent, + +ERC20Component::HasComponent, + +SRC5Component::HasComponent, + +ERC4626HooksTrait, + +Drop + > of IERC4626Camel> { + fn totalSupply(self: @ComponentState) -> u256 { self.total_supply() } - fn balanceOf(self: @ContractState, account: ContractAddress) -> u256 { + fn balanceOf(self: @ComponentState, account: ContractAddress) -> u256 { self.balance_of(account) } fn transferFrom( - ref self: ContractState, + ref self: ComponentState, sender: ContractAddress, recipient: ContractAddress, amount: u256 @@ -244,8 +299,31 @@ mod ERC4626 { } #[generate_trait] - impl InternalImpl of InternalImplTrait { - fn _convert_to_assets(self: @ContractState, shares: u256, round: bool) -> u256 { + pub impl InternalImpl< + TContractState, +HasComponent, + impl erc20: ERC20Component::HasComponent, + impl src5: SRC5Component::HasComponent, + impl Hooks: ERC4626HooksTrait, + +Drop + > of InternalImplTrait { + fn initializer( + ref self: ComponentState, asset: ContractAddress, name: ByteArray, symbol: ByteArray, offset: u8 + ) { + let dispatcher = ERC20ABIDispatcher { contract_address: asset }; + self.offset.write(offset); + let decimals = dispatcher.decimals(); + let mut erc20_comp_mut = get_dep_component_mut!(ref self, erc20); + erc20_comp_mut.initializer(name, symbol); + self.asset.write(asset); + self.underlying_decimals.write(decimals); + + // ! To register interface + // let mut src5_component = get_dep_component_mut!(ref self, src5); + // src5_component.register_interface(interface::IERC721_ID); + // src5_component.register_interface(interface::IERC721_METADATA_ID); + } + + fn _convert_to_assets(self: @ComponentState, shares: u256, round: bool) -> u256 { let total_assets = self.total_assets() + 1; let total_shares = self.total_supply() + pow_256(10, self._decimals_offset()); let assets = shares * total_assets / total_shares; @@ -256,7 +334,7 @@ mod ERC4626 { } } - fn _convert_to_shares(self: @ContractState, assets: u256, round: bool) -> u256 { + fn _convert_to_shares(self: @ComponentState, assets: u256, round: bool) -> u256 { let total_assets = self.total_assets() + 1; let total_shares = self.total_supply() + pow_256(10, self._decimals_offset()); let share = assets * total_shares / total_assets; @@ -268,44 +346,88 @@ mod ERC4626 { } fn _deposit( - ref self: ContractState, + ref self: ComponentState, caller: ContractAddress, receiver: ContractAddress, assets: u256, shares: u256 ) { + Hooks::before_deposit(ref self, caller, receiver, assets, shares); + let dispatcher = ERC20ABIDispatcher { contract_address: self.asset.read() }; - dispatcher.transfer_from(caller, get_contract_address(), assets); - self.erc20._mint(receiver, shares); + dispatcher.transferFrom(caller, get_contract_address(), assets); + let mut erc20_comp_mut = get_dep_component_mut!(ref self, erc20); + erc20_comp_mut.mint(receiver, shares); self.emit(Deposit { sender: caller, owner: receiver, assets, shares }); + + Hooks::after_deposit(ref self, caller, receiver, assets, shares); } fn _withdraw( - ref self: ContractState, + ref self: ComponentState, caller: ContractAddress, receiver: ContractAddress, owner: ContractAddress, assets: u256, shares: u256 ) { + Hooks::before_withdraw(ref self, caller, receiver, owner, assets, shares); + + let mut erc20_comp_mut = get_dep_component_mut!(ref self, erc20); if (caller != owner) { let allowance = self.allowance(owner, caller); - if (allowance != BoundedU256::max()) { + if (allowance != Bounded::::MAX) { assert(allowance >= shares, ERC20Errors::APPROVE_FROM_ZERO); - self.erc20.ERC20_allowances.write((owner, caller), allowance - shares); + erc20_comp_mut.ERC20_allowances.write((owner, caller), allowance - shares); } } - self.erc20._burn(owner, shares); + erc20_comp_mut.burn(owner, shares); let dispatcher = ERC20ABIDispatcher { contract_address: self.asset.read() }; dispatcher.transfer(receiver, assets); self.emit(Withdraw { sender: caller, receiver, owner, assets, shares }); + + Hooks::after_withdraw(ref self, caller, receiver, owner, assets, shares); } - fn _decimals_offset(self: @ContractState) -> u8 { + fn _decimals_offset(self: @ComponentState) -> u8 { self.offset.read() } } } + +pub impl ERC4626HooksEmptyImpl of ERC4626Component::ERC4626HooksTrait { + fn before_deposit( + ref self: ERC4626Component::ComponentState, + caller: ContractAddress, + receiver: ContractAddress, + assets: u256, + shares: u256 + ) {} + fn after_deposit( + ref self: ERC4626Component::ComponentState, + caller: ContractAddress, + receiver: ContractAddress, + assets: u256, + shares: u256 + ) {} + + fn before_withdraw( + ref self: ERC4626Component::ComponentState, + caller: ContractAddress, + receiver: ContractAddress, + owner: ContractAddress, + assets: u256, + shares: u256 + ) {} + fn after_withdraw( + ref self: ERC4626Component::ComponentState, + caller: ContractAddress, + receiver: ContractAddress, + owner: ContractAddress, + assets: u256, + shares: u256 + ) {} +} \ No newline at end of file diff --git a/src/erc4626/interface.cairo b/src/erc4626/interface.cairo index ae77e8a..47748ca 100644 --- a/src/erc4626/interface.cairo +++ b/src/erc4626/interface.cairo @@ -1,12 +1,12 @@ use starknet::ContractAddress; #[starknet::interface] -trait IERC4626 { +pub trait IERC4626 { // ************************************ // * Metadata // ************************************ - fn name(self: @TState) -> felt252; - fn symbol(self: @TState) -> felt252; + fn name(self: @TState) -> ByteArray; + fn symbol(self: @TState) -> ByteArray; fn decimals(self: @TState) -> u8; // ************************************ @@ -63,14 +63,14 @@ trait IERC4626 { #[starknet::interface] -trait IERC4626Metadata { - fn name(self: @TState) -> felt252; - fn symbol(self: @TState) -> felt252; +pub trait IERC4626Metadata { + fn name(self: @TState) -> ByteArray; + fn symbol(self: @TState) -> ByteArray; fn decimals(self: @TState) -> u8; } #[starknet::interface] -trait IERC4626Camel { +pub trait IERC4626Camel { fn totalSupply(self: @TState) -> u256; fn balanceOf(self: @TState, account: ContractAddress) -> u256; fn transferFrom( @@ -79,7 +79,7 @@ trait IERC4626Camel { } #[starknet::interface] -trait IERC4626Snake { +pub trait IERC4626Snake { fn total_supply(self: @TState) -> u256; fn balance_of(self: @TState, account: ContractAddress) -> u256; fn allowance(self: @TState, owner: ContractAddress, spender: ContractAddress) -> u256; @@ -91,7 +91,7 @@ trait IERC4626Snake { } #[starknet::interface] -trait IERC4626Additional { +pub trait IERC4626Additional { fn asset(self: @TState) -> ContractAddress; fn convert_to_assets(self: @TState, shares: u256) -> u256; fn convert_to_shares(self: @TState, assets: u256) -> u256; diff --git a/src/lib.cairo b/src/lib.cairo index ce442a1..7506ffb 100644 --- a/src/lib.cairo +++ b/src/lib.cairo @@ -1,4 +1,11 @@ -mod erc4626; +pub mod erc4626 { + pub mod erc4626; + pub mod interface; +} +mod preset { + pub mod ERC4626; +} + #[cfg(test)] mod tests; mod utils; diff --git a/src/mocks/ERC20.cairo b/src/mocks/ERC20.cairo index 4c8dd7c..85effc3 100644 --- a/src/mocks/ERC20.cairo +++ b/src/mocks/ERC20.cairo @@ -3,7 +3,7 @@ #[starknet::contract] mod ERC20Token { use openzeppelin::access::ownable::OwnableComponent; - use openzeppelin::token::erc20::ERC20Component; + use openzeppelin::token::erc20::{ERC20Component, ERC20HooksEmptyImpl}; use starknet::ContractAddress; use starknet::get_caller_address; @@ -44,21 +44,22 @@ mod ERC20Token { #[constructor] fn constructor(ref self: ContractState, recipient: ContractAddress, initial_supply: u256) { - self.erc20.initializer('Mock', 'MCK'); - self.erc20._mint(recipient, initial_supply); + self.erc20.initializer("Mock", "MCK"); + self.erc20.mint(recipient, initial_supply); } #[generate_trait] - #[external(v0)] impl ExternalImpl of ExternalTrait { + #[abi(per_item)] fn burn(ref self: ContractState, value: u256) { let caller = get_caller_address(); - self.erc20._burn(caller, value); + self.erc20.burn(caller, value); } + #[abi(per_item)] fn mint(ref self: ContractState, recipient: ContractAddress, amount: u256) { self.ownable.assert_only_owner(); - self.erc20._mint(recipient, amount); + self.erc20.mint(recipient, amount); } } } diff --git a/src/preset/ERC4626.cairo b/src/preset/ERC4626.cairo new file mode 100644 index 0000000..eebfb82 --- /dev/null +++ b/src/preset/ERC4626.cairo @@ -0,0 +1,57 @@ +#[starknet::contract] +mod ERC4626 { + use erc4626::erc4626::erc4626::{ERC4626Component, ERC4626HooksEmptyImpl}; + use openzeppelin::introspection::src5::SRC5Component; + use openzeppelin::token::erc20::ERC20Component; + + component!(path: ERC4626Component, storage: erc4626, event: ERC4626Event); + component!(path: ERC20Component, storage: erc20, event: ERC20Event); + component!(path: SRC5Component, storage: src5, event: SRC5Event); + + use starknet::{ContractAddress}; + use openzeppelin::token::erc20::interface::{IERC20, IERC20Dispatcher, IERC20DispatcherTrait}; + use starknet::{get_contract_address}; + + #[abi(embed_v0)] + impl ERC4626AdditionalImpl = ERC4626Component::ERC4626AdditionalImpl; + #[abi(embed_v0)] + impl MetadataEntrypointsImpl = ERC4626Component::MetadataEntrypointsImpl; + #[abi(embed_v0)] + impl SnakeEntrypointsImpl = ERC4626Component::SnakeEntrypointsImpl; + #[abi(embed_v0)] + impl CamelEntrypointsImpl = ERC4626Component::CamelEntrypointsImpl; + + impl ERC4626InternalImpl = ERC4626Component::InternalImpl; + + #[storage] + struct Storage { + #[substorage(v0)] + erc4626: ERC4626Component::Storage, + #[substorage(v0)] + erc20: ERC20Component::Storage, + #[substorage(v0)] + src5: SRC5Component::Storage, + } + + #[event] + #[derive(Drop, starknet::Event)] + enum Event { + #[flat] + ERC4626Event: ERC4626Component::Event, + #[flat] + ERC20Event: ERC20Component::Event, + #[flat] + SRC5Event: SRC5Component::Event, + } + + #[constructor] + fn constructor( + ref self: ContractState, + asset: ContractAddress, + name: ByteArray, + symbol: ByteArray, + offset: u8, + ) { + self.erc4626.initializer(asset, name, symbol, offset); + } +} \ No newline at end of file diff --git a/src/tests/erc4626/test_erc4626.cairo b/src/tests/erc4626/test_erc4626.cairo index 8aa5a6f..53e8f02 100644 --- a/src/tests/erc4626/test_erc4626.cairo +++ b/src/tests/erc4626/test_erc4626.cairo @@ -1,12 +1,13 @@ use core::traits::TryInto; -use debug::PrintTrait; -use erc4626::erc4626::{IERC4626Dispatcher, IERC4626DispatcherTrait}; +use erc4626::erc4626::interface::{IERC4626Dispatcher, IERC4626DispatcherTrait}; use erc4626::utils::{pow_256}; -use integer::BoundedU256; use openzeppelin::token::erc20::{ERC20ABIDispatcher, ERC20ABIDispatcherTrait}; use snforge_std::{ - declare, ContractClassTrait, start_prank, stop_prank, CheatTarget, start_warp, stop_warp + declare, ContractClassTrait, start_cheat_caller_address, stop_cheat_caller_address, + DeclareResultTrait }; +use core::num::traits::Bounded; +use openzeppelin::utils::serde::SerializedAppend; use starknet::{ContractAddress, contract_address_const, get_contract_address}; fn OWNER() -> ContractAddress { @@ -34,25 +35,27 @@ fn VAULT_ADDRESS() -> ContractAddress { } fn deploy_token() -> (ERC20ABIDispatcher, ContractAddress) { - let token = declare('ERC20Token'); + let token = declare("ERC20Token").unwrap().contract_class(); let mut calldata = Default::default(); Serde::serialize(@OWNER(), ref calldata); Serde::serialize(@INITIAL_SUPPLY(), ref calldata); - let address = token.deploy_at(@calldata, TOKEN_ADDRESS()).unwrap(); + let (address, _) = token.deploy_at(@calldata, TOKEN_ADDRESS()).unwrap(); let dispatcher = ERC20ABIDispatcher { contract_address: address, }; (dispatcher, address) } fn deploy_contract() -> (ERC20ABIDispatcher, IERC4626Dispatcher) { let (token, token_address) = deploy_token(); - let mut params = ArrayTrait::::new(); - token_address.serialize(ref params); - params.append('Vault Mock Token'); - params.append('vltMCK'); - params.append(8); - let vault = declare('ERC4626'); - let contract_address = vault.deploy_at(@params, VAULT_ADDRESS()).unwrap(); + let mut calldata = array![]; + let name: ByteArray = "Vault Mock Token"; + let symbol: ByteArray = "vltMCK"; + calldata.append_serde(token_address); + calldata.append_serde(name); + calldata.append_serde(symbol); + calldata.append(0); + let vault = declare("ERC4626").unwrap().contract_class(); + let (contract_address, _) = vault.deploy_at(@calldata, VAULT_ADDRESS()).unwrap(); (token, IERC4626Dispatcher { contract_address }) } @@ -61,70 +64,70 @@ fn deploy_contract() -> (ERC20ABIDispatcher, IERC4626Dispatcher) { fn test_constructor() { let (asset, vault) = deploy_contract(); assert(vault.asset() == asset.contract_address, 'invalid asset'); - assert(vault.decimals() == (18 + 8), 'invalid decimals'); - assert(vault.name() == 'Vault Mock Token', 'invalid name'); - assert(vault.symbol() == 'vltMCK', 'invalid symbol'); + assert(vault.decimals() == (18 + 0), 'invalid decimals'); + assert(vault.name() == "Vault Mock Token", 'invalid name'); + assert(vault.symbol() == "vltMCK", 'invalid symbol'); } #[test] fn convert_to_assets() { - let (asset, vault) = deploy_contract(); - let shares = pow_256(10, 10); + let (_asset, vault) = deploy_contract(); + let shares = pow_256(10, 2); // 10e10 * (0 + 1) / (0 + 10e8) assert(vault.convert_to_assets(shares) == 100, 'invalid assets'); } #[test] fn convert_to_shares() { - let (asset, vault) = deploy_contract(); + let (_asset, vault) = deploy_contract(); let assets = 10; // asset * shares / total assets // 10 * (0 + 10e8) / (0 + 1) - assert(vault.convert_to_shares(assets) == pow_256(10, 9), 'invalid shares'); + assert(vault.convert_to_shares(assets) == pow_256(10, 1), 'invalid shares'); } #[test] fn max_deposit() { - let (asset, vault) = deploy_contract(); - assert(vault.max_deposit(get_contract_address()) == BoundedU256::max(), 'invalid max deposit'); + let (_asset, vault) = deploy_contract(); + assert(vault.max_deposit(get_contract_address()) == Bounded::::MAX, 'invalid max deposit'); } #[test] fn max_mint() { - let (asset, vault) = deploy_contract(); - assert(vault.max_mint(get_contract_address()) == BoundedU256::max(), 'invalid max mint'); + let (_asset, vault) = deploy_contract(); + assert(vault.max_mint(get_contract_address()) == Bounded::::MAX, 'invalid max mint'); } #[test] fn preview_deposit() { - let (asset, vault) = deploy_contract(); - assert(vault.preview_deposit(10) == pow_256(10, 9), 'invalid preview_deposit'); + let (_asset, vault) = deploy_contract(); + assert(vault.preview_deposit(10) == pow_256(10, 1), 'invalid preview_deposit'); } #[test] fn preview_mint() { - let (asset, vault) = deploy_contract(); - assert(vault.preview_mint(pow_256(10, 10)) == 100, 'invalid preview_mint'); + let (_asset, vault) = deploy_contract(); + assert(vault.preview_mint(pow_256(10, 2)) == 100, 'invalid preview_mint'); } #[test] fn preview_redeem() { - let (asset, vault) = deploy_contract(); - assert(vault.preview_redeem(pow_256(10, 10)) == 100, 'invalid preview_redeem'); + let (_asset, vault) = deploy_contract(); + assert(vault.preview_redeem(pow_256(10, 2)) == 100, 'invalid preview_redeem'); } #[test] fn preview_withdraw() { - let (asset, vault) = deploy_contract(); - assert(vault.preview_redeem(pow_256(10, 10)) == 100, 'invalid preview_withdraw'); + let (_asset, vault) = deploy_contract(); + assert(vault.preview_redeem(pow_256(10, 2)) == 100, 'invalid preview_withdraw'); } #[test] fn test_deposit() { let (asset, vault) = deploy_contract(); let amount = asset.balanceOf(OWNER()); - start_prank(CheatTarget::One(asset.contract_address), OWNER()); + start_cheat_caller_address(asset.contract_address, OWNER()); asset.approve(vault.contract_address, amount); - stop_prank(CheatTarget::One(asset.contract_address)); + stop_cheat_caller_address(asset.contract_address); let result = vault.preview_deposit(amount); - start_prank(CheatTarget::One(vault.contract_address), OWNER()); + start_cheat_caller_address(vault.contract_address, OWNER()); assert(vault.deposit(amount, OWNER()) == result, 'invalid shares'); assert(vault.balanceOf(OWNER()) == result, 'invalid balance'); } @@ -133,11 +136,11 @@ fn test_deposit() { fn test_max_redeem() { let (asset, vault) = deploy_contract(); let amount = asset.balanceOf(OWNER()); - start_prank(CheatTarget::One(asset.contract_address), OWNER()); + start_cheat_caller_address(asset.contract_address, OWNER()); asset.approve(vault.contract_address, amount); - stop_prank(CheatTarget::One(asset.contract_address)); - let result = vault.preview_deposit(amount); - start_prank(CheatTarget::One(vault.contract_address), OWNER()); + stop_cheat_caller_address(asset.contract_address); + let _result = vault.preview_deposit(amount); + start_cheat_caller_address(vault.contract_address, OWNER()); let shares = vault.deposit(amount, OWNER()); assert(vault.max_redeem(OWNER()) == shares, 'invalid max redeem'); } @@ -145,12 +148,12 @@ fn test_max_redeem() { fn max_withdraw() { let (asset, vault) = deploy_contract(); let amount = asset.balanceOf(OWNER()); - start_prank(CheatTarget::One(asset.contract_address), OWNER()); + start_cheat_caller_address(asset.contract_address, OWNER()); asset.approve(vault.contract_address, amount); - stop_prank(CheatTarget::One(asset.contract_address)); - let result = vault.preview_deposit(amount); - start_prank(CheatTarget::One(vault.contract_address), OWNER()); - let shares = vault.deposit(amount, OWNER()); + stop_cheat_caller_address(asset.contract_address); + let _result = vault.preview_deposit(amount); + start_cheat_caller_address(vault.contract_address, OWNER()); + let _shares = vault.deposit(amount, OWNER()); let value = vault.convert_to_assets(vault.balanceOf(OWNER())); assert(vault.max_withdraw(OWNER()) == value, 'invalid max withdraw'); } @@ -158,29 +161,29 @@ fn max_withdraw() { fn mint() { let (asset, vault) = deploy_contract(); let amount = asset.balanceOf(OWNER()); - start_prank(CheatTarget::One(asset.contract_address), OWNER()); + start_cheat_caller_address(asset.contract_address, OWNER()); asset.approve(vault.contract_address, amount); - stop_prank(CheatTarget::One(asset.contract_address)); - let result = vault.preview_deposit(amount); + stop_cheat_caller_address(asset.contract_address); + let _result = vault.preview_deposit(amount); let minted = vault.preview_mint(1); - start_prank(CheatTarget::One(vault.contract_address), OWNER()); - let shares = vault.mint(1, OWNER()); + start_cheat_caller_address(vault.contract_address, OWNER()); + let _shares = vault.mint(1, OWNER()); assert(vault.balanceOf(OWNER()) == minted, 'invalid mint shares'); } #[test] fn test_redeem() { let (asset, vault) = deploy_contract(); let amount = asset.balanceOf(OWNER()); - start_prank(CheatTarget::One(asset.contract_address), OWNER()); + start_cheat_caller_address(asset.contract_address, OWNER()); asset.approve(vault.contract_address, amount); - stop_prank(CheatTarget::One(asset.contract_address)); - let result = vault.preview_deposit(amount); - start_prank(CheatTarget::One(vault.contract_address), OWNER()); + stop_cheat_caller_address(asset.contract_address); + let _result = vault.preview_deposit(amount); + start_cheat_caller_address(vault.contract_address, OWNER()); let shares = vault.deposit(amount, OWNER()); assert(vault.balanceOf(OWNER()) == shares, 'invalid balance before'); - let preview = vault.preview_redeem(shares); - start_prank(CheatTarget::One(vault.contract_address), OWNER()); - let redeemed = vault.redeem(shares, OWNER(), OWNER()); + let _preview = vault.preview_redeem(shares); + start_cheat_caller_address(vault.contract_address, OWNER()); + let _redeemed = vault.redeem(shares, OWNER(), OWNER()); assert(vault.balanceOf(OWNER()) == 0, 'invalid balance after'); } @@ -188,16 +191,16 @@ fn test_redeem() { fn test_withdraw() { let (asset, vault) = deploy_contract(); let amount = asset.balanceOf(OWNER()); - start_prank(CheatTarget::One(asset.contract_address), OWNER()); + start_cheat_caller_address(asset.contract_address, OWNER()); asset.approve(vault.contract_address, amount); - stop_prank(CheatTarget::One(asset.contract_address)); - let result = vault.preview_deposit(amount); - start_prank(CheatTarget::One(vault.contract_address), OWNER()); + stop_cheat_caller_address(asset.contract_address); + let _result = vault.preview_deposit(amount); + start_cheat_caller_address(vault.contract_address, OWNER()); let shares = vault.deposit(amount, OWNER()); assert(vault.balanceOf(OWNER()) == shares, 'invalid balance before'); - start_prank(CheatTarget::One(vault.contract_address), OWNER()); - let shares = vault.withdraw(amount, OWNER(), OWNER()); + start_cheat_caller_address(vault.contract_address, OWNER()); + let _shares = vault.withdraw(amount, OWNER(), OWNER()); assert(vault.balanceOf(OWNER()) == 0, 'invalid balance after'); } diff --git a/src/utils.cairo b/src/utils.cairo index 0d0edca..ae776fb 100644 --- a/src/utils.cairo +++ b/src/utils.cairo @@ -1,4 +1,6 @@ -fn pow_256(self: u256, mut exponent: u8) -> u256 { +use core::num::traits::Zero; + +pub fn pow_256(self: u256, mut exponent: u8) -> u256 { if self.is_zero() { return 0; }