diff --git a/contracts/interfaces/IMain.sol b/contracts/interfaces/IMain.sol index f7f2c9bd6..876b4e801 100644 --- a/contracts/interfaces/IMain.sol +++ b/contracts/interfaces/IMain.sol @@ -156,6 +156,8 @@ interface IComponentRegistry { event BrokerSet(IBroker oldVal, IBroker newVal); function broker() external view returns (IBroker); + + function isComponent(address addr) external view returns (bool); } /** diff --git a/contracts/mixins/ComponentRegistry.sol b/contracts/mixins/ComponentRegistry.sol index ff3a29f7c..0bafd742c 100644 --- a/contracts/mixins/ComponentRegistry.sol +++ b/contracts/mixins/ComponentRegistry.sol @@ -35,6 +35,7 @@ abstract contract ComponentRegistry is Initializable, Auth, IComponentRegistry { function _setRToken(IRToken val) private { require(address(val) != address(0), "invalid RToken address"); emit RTokenSet(rToken, val); + isComponent[address(val)] = true; rToken = val; } @@ -43,6 +44,7 @@ abstract contract ComponentRegistry is Initializable, Auth, IComponentRegistry { function _setStRSR(IStRSR val) private { require(address(val) != address(0), "invalid StRSR address"); emit StRSRSet(stRSR, val); + isComponent[address(val)] = true; stRSR = val; } @@ -51,6 +53,7 @@ abstract contract ComponentRegistry is Initializable, Auth, IComponentRegistry { function _setAssetRegistry(IAssetRegistry val) private { require(address(val) != address(0), "invalid AssetRegistry address"); emit AssetRegistrySet(assetRegistry, val); + isComponent[address(val)] = true; assetRegistry = val; } @@ -59,6 +62,7 @@ abstract contract ComponentRegistry is Initializable, Auth, IComponentRegistry { function _setBasketHandler(IBasketHandler val) private { require(address(val) != address(0), "invalid BasketHandler address"); emit BasketHandlerSet(basketHandler, val); + isComponent[address(val)] = true; basketHandler = val; } @@ -67,6 +71,7 @@ abstract contract ComponentRegistry is Initializable, Auth, IComponentRegistry { function _setBackingManager(IBackingManager val) private { require(address(val) != address(0), "invalid BackingManager address"); emit BackingManagerSet(backingManager, val); + isComponent[address(val)] = true; backingManager = val; } @@ -75,6 +80,7 @@ abstract contract ComponentRegistry is Initializable, Auth, IComponentRegistry { function _setDistributor(IDistributor val) private { require(address(val) != address(0), "invalid Distributor address"); emit DistributorSet(distributor, val); + isComponent[address(val)] = true; distributor = val; } @@ -83,6 +89,7 @@ abstract contract ComponentRegistry is Initializable, Auth, IComponentRegistry { function _setRSRTrader(IRevenueTrader val) private { require(address(val) != address(0), "invalid RSRTrader address"); emit RSRTraderSet(rsrTrader, val); + isComponent[address(val)] = true; rsrTrader = val; } @@ -91,6 +98,7 @@ abstract contract ComponentRegistry is Initializable, Auth, IComponentRegistry { function _setRTokenTrader(IRevenueTrader val) private { require(address(val) != address(0), "invalid RTokenTrader address"); emit RTokenTraderSet(rTokenTrader, val); + isComponent[address(val)] = true; rTokenTrader = val; } @@ -99,6 +107,7 @@ abstract contract ComponentRegistry is Initializable, Auth, IComponentRegistry { function _setFurnace(IFurnace val) private { require(address(val) != address(0), "invalid Furnace address"); emit FurnaceSet(furnace, val); + isComponent[address(val)] = true; furnace = val; } @@ -107,13 +116,35 @@ abstract contract ComponentRegistry is Initializable, Auth, IComponentRegistry { function _setBroker(IBroker val) private { require(address(val) != address(0), "invalid Broker address"); emit BrokerSet(broker, val); + isComponent[address(val)] = true; broker = val; } + // 4.1.0 - Required for global lock + mapping(address => bool) public isComponent; + + modifier onlyComponent() { + require(isComponent[_msgSender()], "not a component"); + _; + } + + function cacheComponents() external { + isComponent[address(rToken)] = true; + isComponent[address(stRSR)] = true; + isComponent[address(assetRegistry)] = true; + isComponent[address(basketHandler)] = true; + isComponent[address(backingManager)] = true; + isComponent[address(distributor)] = true; + isComponent[address(rsrTrader)] = true; + isComponent[address(rTokenTrader)] = true; + isComponent[address(furnace)] = true; + isComponent[address(broker)] = true; + } + /** * @dev This empty reserved space is put in place to allow future versions to add new * variables without shifting down storage in the inheritance chain. * See https://docs.openzeppelin.com/contracts/4.x/upgradeable#storage_gaps */ - uint256[40] private __gap; + uint256[39] private __gap; } diff --git a/contracts/p1/Main.sol b/contracts/p1/Main.sol index 418a55d6b..14a2a542d 100644 --- a/contracts/p1/Main.sol +++ b/contracts/p1/Main.sol @@ -161,11 +161,11 @@ contract MainP1 is // === Control Flow === - function beginTx() external virtual { + function beginTx() external virtual onlyComponent { _nonReentrantBefore(); } - function endTx() external virtual { + function endTx() external virtual onlyComponent { _nonReentrantAfter(); } diff --git a/test/Main.test.ts b/test/Main.test.ts index 3405da86a..49ff65dff 100644 --- a/test/Main.test.ts +++ b/test/Main.test.ts @@ -259,6 +259,18 @@ describe(`MainP${IMPLEMENTATION} contract`, () => { expect(await main.rsrTrader()).to.equal(rsrTrader.address) expect(await main.rTokenTrader()).to.equal(rTokenTrader.address) + // Components registered + expect(await main.isComponent(rToken.address)).to.equal(true) + expect(await main.isComponent(stRSR.address)).to.equal(true) + expect(await main.isComponent(assetRegistry.address)).to.equal(true) + expect(await main.isComponent(basketHandler.address)).to.equal(true) + expect(await main.isComponent(backingManager.address)).to.equal(true) + expect(await main.isComponent(distributor.address)).to.equal(true) + expect(await main.isComponent(rsrTrader.address)).to.equal(true) + expect(await main.isComponent(rTokenTrader.address)).to.equal(true) + expect(await main.isComponent(furnace.address)).to.equal(true) + expect(await main.isComponent(broker.address)).to.equal(true) + // Configuration const [rTokenTotal, rsrTotal] = await distributor.totals() expect(rTokenTotal).to.equal(bn(4000)) @@ -3998,6 +4010,50 @@ describe(`MainP${IMPLEMENTATION} contract`, () => { ).to.be.revertedWithCustomError(main, 'ReentrancyGuardReentrantCall') } }) + + it('Should allow to cache components', async () => { + await (await ethers.getContractAt('MainP1', main.address)).cacheComponents() + + expect(await main.isComponent(rToken.address)).to.equal(true) + expect(await main.isComponent(stRSR.address)).to.equal(true) + expect(await main.isComponent(assetRegistry.address)).to.equal(true) + expect(await main.isComponent(basketHandler.address)).to.equal(true) + expect(await main.isComponent(backingManager.address)).to.equal(true) + expect(await main.isComponent(distributor.address)).to.equal(true) + expect(await main.isComponent(rsrTrader.address)).to.equal(true) + expect(await main.isComponent(rTokenTrader.address)).to.equal(true) + expect(await main.isComponent(furnace.address)).to.equal(true) + expect(await main.isComponent(broker.address)).to.equal(true) + }) + + it('Should only allow components to begin-end txs', async () => { + await expect(main.connect(owner).beginTx()).to.be.revertedWith('not a component') + await expect(main.connect(owner).endTx()).to.be.revertedWith('not a component') + await expect(main.connect(other).beginTx()).to.be.revertedWith('not a component') + await expect(main.connect(other).endTx()).to.be.revertedWith('not a component') + + // Try with components + const components: Parameters[0] = { + rToken: rToken.address, + stRSR: stRSR.address, + assetRegistry: assetRegistry.address, + basketHandler: basketHandler.address, + backingManager: backingManager.address, + distributor: distributor.address, + rsrTrader: rsrTrader.address, + rTokenTrader: rTokenTrader.address, + furnace: furnace.address, + broker: broker.address, + } + + // Loop through all components + for (const comp of Object.values(components)) { + await whileImpersonating(comp, async (compSigner) => { + await expect(main.connect(compSigner).beginTx()).to.not.be.reverted + await expect(main.connect(compSigner).endTx()).to.not.be.reverted + }) + } + }) }) describeGas('Gas Reporting', () => {