From d6e72d89b2432ffe9a235131275809e0baeb5075 Mon Sep 17 00:00:00 2001 From: Viktoriia Zotova Date: Wed, 30 Apr 2025 16:28:33 -0400 Subject: [PATCH 1/2] Draft of the method for import ritual --- .../contracts/coordination/Coordinator.sol | 112 ++++++----- contracts/test/CoordinatorTestSet.sol | 34 ---- tests/test_coordinator.py | 179 +++++++++--------- 3 files changed, 148 insertions(+), 177 deletions(-) diff --git a/contracts/contracts/coordination/Coordinator.sol b/contracts/contracts/coordination/Coordinator.sol index 0c8b77d97..92677921e 100644 --- a/contracts/contracts/coordination/Coordinator.sol +++ b/contracts/contracts/coordination/Coordinator.sol @@ -21,6 +21,7 @@ contract Coordinator is Initializable, AccessControlDefaultAdminRulesUpgradeable // DKG Protocol event StartRitual(uint32 indexed ritualId, address indexed authority, address[] participants); + event ImportRitual(uint32 indexed ritualId); event StartAggregationRound(uint32 indexed ritualId); event EndRitual(uint32 indexed ritualId, bool successful); event TranscriptPosted(uint32 indexed ritualId, address indexed node, bytes32 transcriptDigest); @@ -98,21 +99,15 @@ contract Coordinator is Initializable, AccessControlDefaultAdminRulesUpgradeable ITACoChildApplication public immutable application; uint96 private immutable minAuthorization; // TODO use child app for checking eligibility - Ritual[] internal ritualsStub; // former rituals, "internal" for testing only uint32 public timeout; uint16 public maxDkgSize; - bool private stub1; // former isInitiationPublic - - uint256 private stub2; // former totalPendingFees - mapping(uint256 => uint256) private stub3; // former pendingFees - address private stub4; // former feeModel IReimbursementPool internal reimbursementPool; mapping(address => ParticipantKey[]) internal participantKeysHistory; mapping(bytes32 => uint32) internal ritualPublicKeyRegistry; mapping(IFeeModel => bool) public feeModelsRegistry; - mapping(uint256 index => Ritual ritual) internal _rituals; + mapping(uint256 index => Ritual ritual) public rituals; uint256 public numberOfRituals; // Note: Adjust the __preSentinelGap size if more contract variables are added @@ -136,64 +131,15 @@ contract Coordinator is Initializable, AccessControlDefaultAdminRulesUpgradeable __AccessControlDefaultAdminRules_init(0, _admin); } - /// @dev use `upgradeAndCall` for upgrading together with re-initialization - function initializeNumberOfRituals() external reinitializer(2) { - if (numberOfRituals == 0) { - numberOfRituals = ritualsStub.length; - } - } - /// @dev use `upgradeAndCall` for upgrading together with re-initialization function reinitializeDefaultAdmin(address newDefaultAdmin) external reinitializer(3) { _beginDefaultAdminTransfer(newDefaultAdmin); } - function rituals( - uint256 ritualId // uint256 for backward compatibility - ) - external - view - returns ( - address initiator, - uint32 initTimestamp, - uint32 endTimestamp, - uint16 totalTranscripts, - uint16 totalAggregations, - // - address authority, - uint16 dkgSize, - uint16 threshold, - bool aggregationMismatch, - // - IEncryptionAuthorizer accessController, - BLS12381.G1Point memory publicKey, - bytes memory aggregatedTranscript, - IFeeModel feeModel - ) - { - Ritual storage ritual = storageRitual(uint32(ritualId)); - initiator = ritual.initiator; - initTimestamp = ritual.initTimestamp; - endTimestamp = ritual.endTimestamp; - totalTranscripts = ritual.totalTranscripts; - totalAggregations = ritual.totalAggregations; - authority = ritual.authority; - dkgSize = ritual.dkgSize; - threshold = ritual.threshold; - aggregationMismatch = ritual.aggregationMismatch; - accessController = ritual.accessController; - publicKey = ritual.publicKey; - aggregatedTranscript = ritual.aggregatedTranscript; - feeModel = ritual.feeModel; - } - // for backward compatibility function storageRitual(uint32 ritualId) internal view returns (Ritual storage) { - if (ritualId < ritualsStub.length) { - return ritualsStub[ritualId]; - } require(ritualId < numberOfRituals, "Ritual id out of bounds"); - return _rituals[ritualId]; + return rituals[ritualId]; } function getInitiator(uint32 ritualId) external view returns (address) { @@ -359,7 +305,7 @@ contract Coordinator is Initializable, AccessControlDefaultAdminRulesUpgradeable require(duration >= 24 hours, "Invalid ritual duration"); // TODO: Define minimum duration #106 uint32 id = uint32(numberOfRituals); - Ritual storage ritual = _rituals[id]; + Ritual storage ritual = rituals[id]; numberOfRituals += 1; ritual.initiator = msg.sender; ritual.authority = authority; @@ -395,6 +341,56 @@ contract Coordinator is Initializable, AccessControlDefaultAdminRulesUpgradeable return id; } + function importRitual( + uint32 ritualId, + address initiator, + uint32 initTimestamp, + uint32 endTimestamp, + uint16 totalTranscripts, + uint16 totalAggregations, + // + address authority, + uint16 dkgSize, + // uint16 threshold, + // bool aggregationMismatch, + // + // IEncryptionAuthorizer accessController, + BLS12381.G1Point memory publicKey, + bytes memory aggregatedTranscript, + //IFeeModel feeModel, + Participant[] calldata participant + ) external onlyRole(DEFAULT_ADMIN_ROLE) { + require(ritualId == numberOfRituals, "Ritual id out of bounds"); + + Ritual storage ritual = rituals[numberOfRituals]; + numberOfRituals += 1; + ritual.initiator = initiator; + ritual.initTimestamp = initTimestamp; + ritual.endTimestamp = endTimestamp; + ritual.totalTranscripts = totalTranscripts; + ritual.totalAggregations = totalAggregations; + ritual.authority = authority; + ritual.dkgSize = dkgSize; + // ritual.threshold = threshold; + // ritual.aggregationMismatch = aggregationMismatch; + // ritual.accessController = accessController; + ritual.publicKey = publicKey; + ritual.aggregatedTranscript = aggregatedTranscript; + // ritual.feeModel = feeModel; + + for (uint256 i = 0; i < participant.length; i++) { + Participant storage newParticipant = ritual.participant.push(); + Participant calldata current = participant[i]; + newParticipant.provider = current.provider; + newParticipant.aggregated = current.aggregated; + newParticipant.transcript = current.transcript; + newParticipant.decryptionRequestStaticKey = current.decryptionRequestStaticKey; + } + bytes32 registryKey = keccak256(abi.encodePacked(BLS12381.g1PointToBytes(publicKey))); + ritualPublicKeyRegistry[registryKey] = ritualId + 1; + emit ImportRitual(ritualId); + } + function cohortFingerprint(address[] calldata nodes) public pure returns (bytes32) { return keccak256(abi.encode(nodes)); } diff --git a/contracts/test/CoordinatorTestSet.sol b/contracts/test/CoordinatorTestSet.sol index 6b777d116..1d7a9aac3 100644 --- a/contracts/test/CoordinatorTestSet.sol +++ b/contracts/test/CoordinatorTestSet.sol @@ -34,37 +34,3 @@ contract ChildApplicationForCoordinatorMock is ITACoChildApplication { // solhint-disable-next-line no-empty-blocks function penalize(address _stakingProvider) external {} } - -contract ExtendedCoordinator is Coordinator { - constructor(ITACoChildApplication _application) Coordinator(_application) {} - - function initiateOldRitual( - IFeeModel feeModel, - address[] calldata providers, - address authority, - uint32 duration, - IEncryptionAuthorizer accessController - ) external returns (uint32) { - uint16 length = uint16(providers.length); - - uint32 id = uint32(ritualsStub.length); - Ritual storage ritual = ritualsStub.push(); - ritual.initiator = msg.sender; - ritual.authority = authority; - ritual.dkgSize = length; - ritual.threshold = getThresholdForRitualSize(length); - ritual.initTimestamp = uint32(block.timestamp); - ritual.endTimestamp = ritual.initTimestamp + duration; - ritual.accessController = accessController; - ritual.feeModel = feeModel; - - address previous = address(0); - for (uint256 i = 0; i < length; i++) { - Participant storage newParticipant = ritual.participant.push(); - address current = providers[i]; - newParticipant.provider = current; - previous = current; - } - return id; - } -} diff --git a/tests/test_coordinator.py b/tests/test_coordinator.py index 104930876..d77039807 100644 --- a/tests/test_coordinator.py +++ b/tests/test_coordinator.py @@ -2,7 +2,6 @@ import ape import pytest -from ape.utils import ZERO_ADDRESS from eth_account import Account from hexbytes import HexBytes from web3 import Web3 @@ -60,7 +59,7 @@ def erc20(project, initiator): @pytest.fixture() def coordinator(project, deployer, application, oz_dependency): admin = deployer - contract = project.ExtendedCoordinator.deploy( + contract = project.Coordinator.deploy( application.address, sender=deployer, ) @@ -72,7 +71,7 @@ def coordinator(project, deployer, application, oz_dependency): encoded_initializer_function, sender=deployer, ) - proxy_contract = project.ExtendedCoordinator.at(proxy.address) + proxy_contract = project.Coordinator.at(proxy.address) return proxy_contract @@ -597,88 +596,6 @@ def test_post_aggregation_fails( fee_model.withdrawTokens(fee_model_balance_after_refund, sender=deployer) -def test_upgrade( - coordinator, nodes, initiator, erc20, fee_model, treasury, deployer, global_allow_list -): - coordinator.initiateOldRitual( - fee_model, nodes, initiator, DURATION, global_allow_list.address, sender=initiator - ) - coordinator.initiateOldRitual( - ZERO_ADDRESS, [nodes[0]], treasury, DURATION // 2, deployer, sender=initiator - ) - assert coordinator.numberOfRituals() == 0 - coordinator.initializeNumberOfRituals(sender=deployer) - assert coordinator.numberOfRituals() == 2 - - initiate_ritual( - coordinator=coordinator, - fee_model=fee_model, - erc20=erc20, - authority=initiator, - nodes=nodes, - allow_logic=global_allow_list, - ) - assert coordinator.numberOfRituals() == 3 - - assert coordinator.getRitualState(0) == RitualState.DKG_AWAITING_TRANSCRIPTS - assert coordinator.getRitualState(1) == RitualState.DKG_AWAITING_TRANSCRIPTS - assert coordinator.getRitualState(2) == RitualState.DKG_AWAITING_TRANSCRIPTS - - ritual_struct = coordinator.rituals(0) - assert ritual_struct["initiator"] == initiator - init, end = ritual_struct["initTimestamp"], ritual_struct["endTimestamp"] - assert end - init == DURATION - total_transcripts, total_aggregations = ( - ritual_struct["totalTranscripts"], - ritual_struct["totalAggregations"], - ) - assert total_transcripts == total_aggregations == 0 - assert ritual_struct["authority"] == initiator - assert ritual_struct["dkgSize"] == len(nodes) - assert ritual_struct["threshold"] == 1 + len(nodes) // 2 - assert not ritual_struct["aggregationMismatch"] - assert ritual_struct["accessController"] == global_allow_list.address - assert ritual_struct["publicKey"] == (b"\x00" * 32, b"\x00" * 16) - assert not ritual_struct["aggregatedTranscript"] - assert ritual_struct["feeModel"] == fee_model.address - - ritual_struct = coordinator.rituals(1) - assert ritual_struct["initiator"] == initiator - init, end = ritual_struct["initTimestamp"], ritual_struct["endTimestamp"] - assert end - init == DURATION // 2 - total_transcripts, total_aggregations = ( - ritual_struct["totalTranscripts"], - ritual_struct["totalAggregations"], - ) - assert total_transcripts == total_aggregations == 0 - assert ritual_struct["authority"] == treasury - assert ritual_struct["dkgSize"] == 1 - assert ritual_struct["threshold"] == 1 # threshold - assert not ritual_struct["aggregationMismatch"] # aggregationMismatch - assert ritual_struct["accessController"] == deployer # accessController - assert ritual_struct["publicKey"] == (b"\x00" * 32, b"\x00" * 16) # publicKey - assert not ritual_struct["aggregatedTranscript"] # aggregatedTranscript - assert ritual_struct["feeModel"] == ZERO_ADDRESS # feeModel - - ritual_struct = coordinator.rituals(2) - assert ritual_struct["initiator"] == initiator - init, end = ritual_struct["initTimestamp"], ritual_struct["endTimestamp"] - assert end - init == DURATION - total_transcripts, total_aggregations = ( - ritual_struct["totalTranscripts"], - ritual_struct["totalAggregations"], - ) - assert total_transcripts == total_aggregations == 0 - assert ritual_struct["authority"] == initiator - assert ritual_struct["dkgSize"] == len(nodes) - assert ritual_struct["threshold"] == 1 + len(nodes) // 2 # threshold - assert not ritual_struct["aggregationMismatch"] # aggregationMismatch - assert ritual_struct["accessController"] == global_allow_list.address # accessController - assert ritual_struct["publicKey"] == (b"\x00" * 32, b"\x00" * 16) # publicKey - assert not ritual_struct["aggregatedTranscript"] # aggregatedTranscript - assert ritual_struct["feeModel"] == fee_model.address # feeModel - - def test_withdraw_tokens(coordinator, initiator, erc20, treasury, deployer): # Let's send some tokens to Coordinator by mistake erc20.transfer(coordinator.address, 42, sender=initiator) @@ -712,3 +629,95 @@ def test_transfer_ownership( coordinator.reinitializeDefaultAdmin(treasury.address, sender=deployer) coordinator.acceptDefaultAdminTransfer(sender=treasury) assert coordinator.defaultAdmin() == treasury.address + + +def test_import_ritual( + coordinator, nodes, initiator, erc20, fee_model, deployer, treasury, global_allow_list +): + authority, tx = initiate_ritual( + coordinator=coordinator, + fee_model=fee_model, + erc20=erc20, + authority=initiator, + nodes=nodes, + allow_logic=global_allow_list, + ) + ritual_id = 0 + + size = len(nodes) + threshold = coordinator.getThresholdForRitualSize(size) + transcript = generate_transcript(size, threshold) + + for node in nodes: + coordinator.postTranscript(ritual_id, transcript, sender=node) + + aggregated = transcript # has the same size as transcript + decryption_request_static_keys = [os.urandom(42) for _ in nodes] + dkg_public_key = (os.urandom(32), os.urandom(16)) + for i, node in enumerate(nodes): + coordinator.postAggregation( + ritual_id, aggregated, dkg_public_key, decryption_request_static_keys[i], sender=node + ) + + ritual_id = coordinator.numberOfRituals() + ritual = coordinator.rituals(ritual_id - 1) + + participants = [] + for i, node in enumerate(nodes): + participants.append((nodes[i], True, transcript, decryption_request_static_keys[i])) + + tx = coordinator.importRitual( + ritual_id, + ritual.initiator, + ritual.initTimestamp, + ritual.endTimestamp, + ritual.totalTranscripts, + ritual.totalAggregations, + ritual.authority, + ritual.dkgSize, + # ritual.threshold, + # ritual.aggregationMismatch, + # ritual.accessController, + ritual.publicKey, + ritual.aggregatedTranscript, + # ritual.feeModel, + participants, + sender=deployer, + ) + + events = [event for event in tx.events if event.event_name == "ImportRitual"] + assert len(events) == 1 + event = events[0] + assert event.ritualId == ritual_id + + assert coordinator.getRitualState(ritual_id) == RitualState.ACTIVE + + ritual_struct = coordinator.rituals(ritual_id) + assert ritual_struct.initiator == initiator + init, end = ritual_struct.initTimestamp, ritual_struct.endTimestamp + assert end - init == DURATION + total_transcripts, total_aggregations = ( + ritual_struct.totalTranscripts, + ritual_struct.totalAggregations, + ) + assert total_transcripts == total_aggregations == size + assert ritual_struct.authority == authority + assert ritual_struct.dkgSize == size + # assert ritual_struct.threshold == 1 + size // 2 # threshold + assert not ritual_struct.aggregationMismatch # aggregationMismatch + # assert ritual_struct.accessController == global_allow_list.address # accessController + # assert ritual_struct.feeModel == fee_model.address # feeModel + assert ritual_struct.publicKey == dkg_public_key # publicKey + assert ritual_struct.aggregatedTranscript == aggregated # aggregatedTranscript + + participants = coordinator.getParticipants(ritual_id) + assert len(participants) == size + for i, participant in enumerate(participants): + assert participant.provider == nodes[i] + assert participant.aggregated + assert participant.transcript == transcript + assert participant.decryptionRequestStaticKey == decryption_request_static_keys[i] + + retrieved_public_key = coordinator.getPublicKeyFromRitualId(ritual_id) + assert retrieved_public_key == dkg_public_key + assert coordinator.getRitualIdFromPublicKey(dkg_public_key) == ritual_id From 6bf45e59c2cd74406c658137238e4f0e4215f9df Mon Sep 17 00:00:00 2001 From: Viktoriia Zotova Date: Wed, 7 May 2025 14:28:56 -0400 Subject: [PATCH 2/2] Draft of script to migrate rituals --- scripts/export_ritual.py | 103 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 103 insertions(+) create mode 100644 scripts/export_ritual.py diff --git a/scripts/export_ritual.py b/scripts/export_ritual.py new file mode 100644 index 000000000..896a129f9 --- /dev/null +++ b/scripts/export_ritual.py @@ -0,0 +1,103 @@ +#!/usr/bin/python3 + +import click +from ape import networks +from ape.cli import ConnectedProviderCommand, account_option + +from deployment import registry +from deployment.constants import ACCESS_CONTROLLERS, SUPPORTED_TACO_DOMAINS +from deployment.params import Transactor +from deployment.types import ChecksumAddress +from deployment.utils import check_plugins + + +@click.command(cls=ConnectedProviderCommand, name="export-rituals") +@account_option() +# @network_option(required=True) +@click.option( + "--domain-from", + "-df", + help="From TACo domain", + type=click.Choice(SUPPORTED_TACO_DOMAINS), + required=True, +) +@click.option( + "--domain-to", + "-dt", + help="To TACo domain", + type=click.Choice(SUPPORTED_TACO_DOMAINS), + required=True, +) +@click.option( + "--access-controller", + "-c", + help="The registry name of an access controller contract.", + type=click.Choice(ACCESS_CONTROLLERS), + required=True, +) +@click.option( + "--fee-model", + "-f", + help="The address of the fee model/subscription contract.", + type=ChecksumAddress(), + required=True, +) +@click.option("--ritual-id", "-r", help="Ritual ID to check", type=int, required=True) +def cli( + domain_from, + domain_to, + account, + # network, + access_controller, + fee_model, + ritual_id, +): + """Export a ritual between TACo domains.""" + + # Setup + check_plugins() + # click.echo(f"Connected to {network.name} network.") + + # Get the contracts from the registry + ritual = None + participants = None + with networks.polygon.mainnet.use_provider("infura"): + coordinator_contract_from = registry.get_contract( + domain=domain_from, contract_name="Coordinator" + ) + ritual = coordinator_contract_from.rituals(ritual_id) + participants = coordinator_contract_from.getParticipants(ritual_id) + + with networks.polygon.sepolia.use_provider("infura"): + coordinator_contract_to = registry.get_contract( + domain=domain_to, contract_name="Coordinator" + ) + # access_controller_contract = registry.get_contract( + # domain=domain_to, contract_name=access_controller + # ) + # fee_model_contract = Contract(fee_model) + + # Initiate the ritual + transactor = Transactor(account=account) + transactor.transact( + coordinator_contract_to.importRitual, + ritual_id, + ritual.initiator, + ritual.initTimestamp, + ritual.endTimestamp, + ritual.totalTranscripts, + ritual.totalAggregations, + ritual.authority, + ritual.dkgSize, + # ritual.threshold, + # ritual.aggregationMismatch, + # access_controller_contract.address, + ritual.publicKey, + ritual.aggregatedTranscript, + # fee_model_contract.address, + participants, + ) + + +if __name__ == "__main__": + cli()