diff --git a/.gitignore b/.gitignore index ea7a40e..540925f 100644 --- a/.gitignore +++ b/.gitignore @@ -21,7 +21,6 @@ dist/ downloads/ eggs/ .eggs/ -lib/ lib64/ parts/ sdist/ diff --git a/.gitmodules b/.gitmodules new file mode 100644 index 0000000..92c633f --- /dev/null +++ b/.gitmodules @@ -0,0 +1,4 @@ +[submodule "cairo_contracts"] + url = https://github.com/OpenZeppelin/cairo-contracts + path = lib/cairo_contracts + branch = refs/heads/0.1.0 diff --git a/contracts/interfaces/IAccessController.cairo b/contracts/interfaces/IAccessController.cairo new file mode 100644 index 0000000..2116f1e --- /dev/null +++ b/contracts/interfaces/IAccessController.cairo @@ -0,0 +1,30 @@ +%lang starknet + +from starkware.cairo.common.uint256 import Uint256 + +@contract_interface +namespace IAccessController: + func isAllowed(address: felt) -> (is_allowed: felt): + end + + func freeSlotsCount() -> (free_slots_count: felt): + end + + func getOwner() -> (owner: felt): + end + + func increaseMaxSlots(increase_max_slots_by: felt) -> (): + end + + func register() -> (): + end + + func forceRegister(address: felt) -> (): + end + + func forceRegisterBatch(batch_address_len : felt, batch_address : felt*) -> (): + end + + func transferOwnership(new_owner: felt) -> (new_owner: felt): + end +end diff --git a/lib/cairo_contracts b/lib/cairo_contracts new file mode 160000 index 0000000..fd6630e --- /dev/null +++ b/lib/cairo_contracts @@ -0,0 +1 @@ +Subproject commit fd6630e327d73cdfeb8660d56768c1697517b44d diff --git a/protostar.toml b/protostar.toml new file mode 100644 index 0000000..cf78856 --- /dev/null +++ b/protostar.toml @@ -0,0 +1,13 @@ +["protostar.config"] +protostar_version = "0.2.3" + +["protostar.project"] +libs_path = "lib" + +["protostar.contracts"] +AccessController = [ + "./contracts/AccessController.cairo", +] + +["protostar.shared_command_configs"] +cairo_path = ["./lib/cairo_contracts/src"] \ No newline at end of file diff --git a/tests/requirements.txt b/tests/requirements.txt index 09565ac..995c6f4 100644 --- a/tests/requirements.txt +++ b/tests/requirements.txt @@ -1,5 +1,5 @@ flake8==4.0.1 pytest==7.1.2 pytest-asyncio==0.18.3 -cairo-lang==0.8.1 +cairo-lang==0.9.0 pytest-describe==2.0.1 diff --git a/tests/test_AccessController.cairo b/tests/test_AccessController.cairo new file mode 100644 index 0000000..f9fe558 --- /dev/null +++ b/tests/test_AccessController.cairo @@ -0,0 +1,187 @@ +%lang starknet + +from contracts.interfaces.IAccessController import IAccessController +from starkware.starknet.common.syscalls import get_caller_address + + +@view +func __setup__(): + %{ + context.owner = 123 + context.contract_address = deploy_contract("./contracts/AccessController.cairo", [10, context.owner]).contract_address + %} + return () +end + +@external +func test_is_allowed_should_return_false_when_address_is_not_in_whitelist{syscall_ptr : felt*, range_check_ptr}(): + # Given + tempvar contract_address + tempvar owner + %{ ids.contract_address = context.contract_address %} + %{ ids.owner = context.owner %} + + # When + let (is_allowed) = IAccessController.isAllowed(contract_address=contract_address, address=owner) + + # Then + assert is_allowed = 0 + return () +end + +@external +func test_free_slot_count_should_return_remaining_slot_number{syscall_ptr : felt*, range_check_ptr}(): + # Given + tempvar contract_address + %{ ids.contract_address = context.contract_address %} + + # When + let (free_slots_count) = IAccessController.freeSlotsCount(contract_address=contract_address) + + # Then + assert free_slots_count = 10 + return () +end + +@external +func test_get_owner_should_return_owner_public_key{syscall_ptr : felt*, range_check_ptr}(): + # Given + tempvar contract_address + tempvar initial_owner + %{ ids.contract_address = context.contract_address %} + %{ ids.initial_owner = context.owner %} + + # When + let (owner) = IAccessController.getOwner(contract_address=contract_address) + + # Then + assert owner = initial_owner + return () +end + +@external +func test_increase_max_slot_should_succeed_when_caller_is_owner{syscall_ptr : felt*, range_check_ptr}(): + # Given + tempvar contract_address + %{ ids.contract_address = context.contract_address %} + %{ stop_prank_callable = start_prank(context.owner, context.contract_address) %} + + # When + IAccessController.increaseMaxSlots(contract_address=contract_address, increase_max_slots_by=2) + + # Then + let (free_slots_count) = IAccessController.freeSlotsCount(contract_address=contract_address) + assert free_slots_count = 12 + %{ stop_prank_callable() %} + return () +end + +@external +func test_increase_max_slot_should_fail_when_caller_is_not_owner{syscall_ptr : felt*, range_check_ptr}(): + # Given + tempvar contract_address + tempvar initial_owner + %{ ids.contract_address = context.contract_address %} + %{ ids.initial_owner = context.owner %} + + # When + %{ expect_revert("TRANSACTION_FAILED") %} + IAccessController.increaseMaxSlots(contract_address=contract_address, increase_max_slots_by=2) + + # Then + let (free_slots_count) = IAccessController.freeSlotsCount(contract_address=contract_address) + assert free_slots_count = 10 + return () +end + +@external +func test_register_should_add_address_to_whitelist{syscall_ptr : felt*, range_check_ptr}(): + # Given + tempvar contract_address + %{ ids.contract_address = context.contract_address %} + %{ stop_prank_callable = start_prank(12345, context.contract_address) %} + + # When + IAccessController.register(contract_address=contract_address) + + # Then + let (is_allowed) = IAccessController.isAllowed(contract_address=contract_address, address=12345) + assert is_allowed = 1 + let (free_slots_count) = IAccessController.freeSlotsCount(contract_address=contract_address) + assert free_slots_count = 9 + %{ stop_prank_callable() %} + return () +end + +@external +func test_force_register_should_add_address_to_whitelist_when_caller_is_owner{syscall_ptr : felt*, range_check_ptr}(): + # Given + tempvar contract_address + %{ ids.contract_address = context.contract_address %} + %{ stop_prank_callable = start_prank(123, context.contract_address) %} + + # When + IAccessController.forceRegister(contract_address=contract_address, address=12345) + + # Then + let (is_allowed) = IAccessController.isAllowed(contract_address=contract_address, address=12345) + assert is_allowed = 1 + let (free_slots_count) = IAccessController.freeSlotsCount(contract_address=contract_address) + assert free_slots_count = 9 + %{ stop_prank_callable() %} + return () +end + +@external +func test_force_register_should_fail_when_caller_is_not_owner{syscall_ptr : felt*, range_check_ptr}(): + # Given + tempvar contract_address + %{ ids.contract_address = context.contract_address %} + + # When + %{ expect_revert("TRANSACTION_FAILED") %} + IAccessController.forceRegister(contract_address=contract_address, address=12345) + + # Then + let (is_allowed) = IAccessController.isAllowed(contract_address=contract_address, address=12345) + assert is_allowed = 0 + let (free_slots_count) = IAccessController.freeSlotsCount(contract_address=contract_address) + assert free_slots_count = 10 + %{ stop_prank_callable() %} + return () +end + +@external +func test_transfer_ownership_should_succeed_when_caller_is_owner{syscall_ptr : felt*, range_check_ptr}(): + # Given + tempvar contract_address + %{ ids.contract_address = context.contract_address %} + %{ stop_prank_callable = start_prank(123, context.contract_address) %} + + # When + IAccessController.transferOwnership(contract_address=contract_address, new_owner=12345) + + # Then + let (owner) = IAccessController.getOwner(contract_address=contract_address) + assert owner = 12345 + %{ stop_prank_callable() %} + return () +end + +@external +func test_transfer_ownership_should_fail_when_caller_is_not_owner{syscall_ptr : felt*, range_check_ptr}(): + # Given + tempvar contract_address + %{ ids.contract_address = context.contract_address %} + %{ stop_prank_callable = start_prank(123, context.contract_address) %} + + # When + %{ expect_revert("TRANSACTION_FAILED") %} + IAccessController.transferOwnership(contract_address=contract_address, new_owner=12345) + + # Then + let (owner) = IAccessController.getOwner(contract_address=contract_address) + assert owner = 123 + %{ stop_prank_callable() %} + return () +end \ No newline at end of file