Files
2026-08-17 17:10:35 -06:00

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;
}
}