105 lines
3.4 KiB
Solidity
105 lines
3.4 KiB
Solidity
// SPDX-License-Identifier: MIT
|
|
pragma solidity 0.8.35;
|
|
|
|
import {ERC20} from "@openzeppelin/contracts/token/ERC20/ERC20.sol";
|
|
|
|
interface IReentrantBankTarget {
|
|
function deposit(uint256 amount) external;
|
|
function withdraw(uint256 amount) external;
|
|
function balanceOf(address account) external view returns (uint256);
|
|
function totalLiabilities() external view returns (uint256);
|
|
}
|
|
|
|
contract ReentrantToken is ERC20 {
|
|
enum Callback {
|
|
None,
|
|
Deposit,
|
|
Withdraw
|
|
}
|
|
|
|
IReentrantBankTarget public callbackTarget;
|
|
Callback public callback;
|
|
bool public propagateRevert;
|
|
bool public nestedCallAttempted;
|
|
bool public nestedCallSucceeded;
|
|
bytes4 public nestedRevertSelector;
|
|
uint256 public observedAccountBalance;
|
|
uint256 public observedLiabilities;
|
|
|
|
constructor() ERC20("Reentrant Token", "REENT") {}
|
|
|
|
function decimals() public pure override returns (uint8) {
|
|
return 6;
|
|
}
|
|
|
|
function mint(address to, uint256 amount) external {
|
|
_mint(to, amount);
|
|
}
|
|
|
|
function configureDepositCallback(address bank, bool propagate) external {
|
|
callbackTarget = IReentrantBankTarget(bank);
|
|
callback = Callback.Deposit;
|
|
propagateRevert = propagate;
|
|
_resetObservations();
|
|
}
|
|
|
|
function configureWithdrawalCallback(address bank, bool propagate) external {
|
|
callbackTarget = IReentrantBankTarget(bank);
|
|
callback = Callback.Withdraw;
|
|
propagateRevert = propagate;
|
|
_resetObservations();
|
|
}
|
|
|
|
function clearCallback() external {
|
|
callback = Callback.None;
|
|
propagateRevert = false;
|
|
_resetObservations();
|
|
}
|
|
|
|
function transferFrom(address from, address to, uint256 amount) public override returns (bool) {
|
|
if (callback == Callback.Deposit && _msgSender() == address(callbackTarget)) {
|
|
observedAccountBalance = callbackTarget.balanceOf(from);
|
|
observedLiabilities = callbackTarget.totalLiabilities();
|
|
_attemptNestedCall(abi.encodeCall(IReentrantBankTarget.deposit, (1)));
|
|
}
|
|
return super.transferFrom(from, to, amount);
|
|
}
|
|
|
|
function transfer(address to, uint256 amount) public override returns (bool) {
|
|
if (callback == Callback.Withdraw && _msgSender() == address(callbackTarget)) {
|
|
observedAccountBalance = callbackTarget.balanceOf(to);
|
|
observedLiabilities = callbackTarget.totalLiabilities();
|
|
_attemptNestedCall(abi.encodeCall(IReentrantBankTarget.withdraw, (1)));
|
|
}
|
|
return super.transfer(to, amount);
|
|
}
|
|
|
|
function _attemptNestedCall(bytes memory callData) private {
|
|
nestedCallAttempted = true;
|
|
bytes memory revertData;
|
|
(nestedCallSucceeded, revertData) = address(callbackTarget).call(callData);
|
|
|
|
if (!nestedCallSucceeded && revertData.length >= 4) {
|
|
bytes4 selector;
|
|
assembly ("memory-safe") {
|
|
selector := mload(add(revertData, 0x20))
|
|
}
|
|
nestedRevertSelector = selector;
|
|
}
|
|
|
|
if (!nestedCallSucceeded && propagateRevert) {
|
|
assembly ("memory-safe") {
|
|
revert(add(revertData, 0x20), mload(revertData))
|
|
}
|
|
}
|
|
}
|
|
|
|
function _resetObservations() private {
|
|
nestedCallAttempted = false;
|
|
nestedCallSucceeded = false;
|
|
nestedRevertSelector = bytes4(0);
|
|
observedAccountBalance = 0;
|
|
observedLiabilities = 0;
|
|
}
|
|
}
|