Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,6 @@ dist/
downloads/
eggs/
.eggs/
lib/
lib64/
parts/
sdist/
Expand Down
4 changes: 4 additions & 0 deletions .gitmodules
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
[submodule "cairo_contracts"]
url = https://github.com/OpenZeppelin/cairo-contracts
path = lib/cairo_contracts
branch = refs/heads/0.1.0
30 changes: 30 additions & 0 deletions contracts/interfaces/IAccessController.cairo
Original file line number Diff line number Diff line change
@@ -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
1 change: 1 addition & 0 deletions lib/cairo_contracts
Submodule cairo_contracts added at fd6630
13 changes: 13 additions & 0 deletions protostar.toml
Original file line number Diff line number Diff line change
@@ -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"]
2 changes: 1 addition & 1 deletion tests/requirements.txt
Original file line number Diff line number Diff line change
@@ -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
187 changes: 187 additions & 0 deletions tests/test_AccessController.cairo
Original file line number Diff line number Diff line change
@@ -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