Skip to content

Commit 8079cda

Browse files
committed
fix: role manager with matching_func
Signed-off-by: Andreas Bichinger <andreas.bichinger@gmail.com>
1 parent 5e16bff commit 8079cda

3 files changed

Lines changed: 265 additions & 16 deletions

File tree

casbin/rbac/default_role_manager/role_manager.py

Lines changed: 31 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -29,11 +29,6 @@ def create_role(self, name):
2929
if name not in self.all_roles.keys():
3030
self.all_roles[name] = Role(name)
3131

32-
if self.matching_func is not None:
33-
for key, role in self.all_roles.items():
34-
if self.matching_func(name, key) and name != key:
35-
self.all_roles[name].add_role(role)
36-
3732
return self.all_roles[name]
3833

3934
def clear(self):
@@ -50,6 +45,17 @@ def add_link(self, name1, name2, *domain):
5045
role2 = self.create_role(name2)
5146
role1.add_role(role2)
5247

48+
if self.matching_func is not None:
49+
for key, role in self.all_roles.items():
50+
if self.matching_func(key, name1) and name1 != key:
51+
self.all_roles[key].add_role(role1)
52+
if self.matching_func(key, name2) and name2 != key:
53+
self.all_roles[name2].add_role(role)
54+
if self.matching_func(name1, key) and name1 != key:
55+
self.all_roles[key].add_role(role1)
56+
if self.matching_func(name2, key) and name2 != key:
57+
self.all_roles[name2].add_role(role)
58+
5359
def delete_link(self, name1, name2, *domain):
5460
if len(domain) == 1:
5561
name1 = domain[0] + "::" + name1
@@ -77,9 +83,14 @@ def has_link(self, name1, name2, *domain):
7783
if not self.has_role(name1) or not self.has_role(name2):
7884
return False
7985

80-
role1 = self.create_role(name1)
81-
82-
return role1.has_role(name2, self.max_hierarchy_level)
86+
if self.matching_func is None:
87+
role1 = self.create_role(name1)
88+
return role1.has_role(name2, self.max_hierarchy_level)
89+
else:
90+
for key, role in self.all_roles.items():
91+
if self.matching_func(name1, key) and role.has_role(name2, self.max_hierarchy_level, self.matching_func):
92+
return True
93+
return False
8394

8495
def get_roles(self, name, *domain):
8596
"""
@@ -158,23 +169,27 @@ def delete_role(self, role):
158169
self.roles.remove(rr)
159170
return
160171

161-
def has_role(self, name, hierarchy_level):
162-
if name == self.name:
172+
def has_role(self, name, hierarchy_level, matching_func=None):
173+
if self.has_direct_role(name, matching_func):
163174
return True
164175
if hierarchy_level <= 0:
165176
return False
166177

167178
for role in self.roles:
168-
if role.has_role(name, hierarchy_level - 1):
179+
if role.has_role(name, hierarchy_level - 1, matching_func):
169180
return True
170181

171182
return False
172183

173-
def has_direct_role(self, name):
174-
for role in self.roles:
175-
if role.name == name:
176-
return True
177-
184+
def has_direct_role(self, name, matching_func=None):
185+
if matching_func is None:
186+
for role in self.roles:
187+
if role.name == name:
188+
return True
189+
else:
190+
for role in self.roles:
191+
if matching_func(name, role.name):
192+
return True
178193
return False
179194

180195
def to_string(self):

tests/rbac/__init__.py

Whitespace-only changes.

tests/rbac/test_role_manager.py

Lines changed: 234 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,234 @@
1+
from unittest import TestCase
2+
from casbin.rbac import default_role_manager
3+
from casbin.util import regex_match_func
4+
import time
5+
from concurrent.futures import ThreadPoolExecutor
6+
7+
def get_role_manager():
8+
return default_role_manager.RoleManager(max_hierarchy_level=10)
9+
10+
class TestDefaultRoleManager(TestCase):
11+
12+
def test_role(self):
13+
rm = get_role_manager()
14+
rm.add_link("u1", "g1")
15+
rm.add_link("u2", "g1")
16+
rm.add_link("u3", "g2")
17+
rm.add_link("u4", "g2")
18+
rm.add_link("u4", "g3")
19+
rm.add_link("g1", "g3")
20+
21+
# Current role inheritance tree:
22+
# g3 g2
23+
# / \ / \
24+
# g1 u4 u3
25+
# / \
26+
# u1 u2
27+
28+
self.assertTrue(rm.has_link("u1", "g1"))
29+
self.assertFalse(rm.has_link("u1", "g2"))
30+
self.assertTrue(rm.has_link("u1", "g3"))
31+
self.assertTrue(rm.has_link("u2", "g1"))
32+
self.assertFalse(rm.has_link("u2", "g2"))
33+
self.assertTrue(rm.has_link("u2", "g3"))
34+
self.assertFalse(rm.has_link("u3", "g1"))
35+
self.assertTrue(rm.has_link("u3", "g2"))
36+
self.assertFalse(rm.has_link("u3", "g3"))
37+
self.assertFalse(rm.has_link("u4", "g1"))
38+
self.assertTrue(rm.has_link("u4", "g2"))
39+
self.assertTrue(rm.has_link("u4", "g3"))
40+
41+
self.assertCountEqual(rm.get_roles("u1"), ["g1"])
42+
self.assertCountEqual(rm.get_roles("u2"), ["g1"])
43+
self.assertCountEqual(rm.get_roles("u3"), ["g2"])
44+
self.assertCountEqual(rm.get_roles("u4"), ["g2", "g3"])
45+
self.assertCountEqual(rm.get_roles("g1"), ["g3"])
46+
self.assertCountEqual(rm.get_roles("g2"), [])
47+
self.assertCountEqual(rm.get_roles("g3"), [])
48+
49+
rm.delete_link("g1", "g3")
50+
rm.delete_link("u4", "g2")
51+
52+
# Current role inheritance tree after deleting the links:
53+
# g3 g2
54+
# \ \
55+
# g1 u4 u3
56+
# / \
57+
# u1 u2
58+
59+
self.assertTrue(rm.has_link("u1", "g1"))
60+
self.assertFalse(rm.has_link("u1", "g2"))
61+
self.assertFalse(rm.has_link("u1", "g3"))
62+
self.assertTrue(rm.has_link("u2", "g1"))
63+
self.assertFalse(rm.has_link("u2", "g2"))
64+
self.assertFalse(rm.has_link("u2", "g3"))
65+
self.assertFalse(rm.has_link("u3", "g1"))
66+
self.assertTrue(rm.has_link("u3", "g2"))
67+
self.assertFalse(rm.has_link("u3", "g3"))
68+
self.assertFalse(rm.has_link("u4", "g1"))
69+
self.assertFalse(rm.has_link("u4", "g2"))
70+
self.assertTrue(rm.has_link("u4", "g3"))
71+
72+
self.assertCountEqual(rm.get_roles("u1"), ["g1"])
73+
self.assertCountEqual(rm.get_roles("u2"), ["g1"])
74+
self.assertCountEqual(rm.get_roles("u3"), ["g2"])
75+
self.assertCountEqual(rm.get_roles("u4"), ["g3"])
76+
self.assertCountEqual(rm.get_roles("g1"), [])
77+
self.assertCountEqual(rm.get_roles("g2"), [])
78+
self.assertCountEqual(rm.get_roles("g3"), [])
79+
80+
def test_domain_role(self):
81+
rm = get_role_manager()
82+
rm.add_link("u1", "g1", "domain1")
83+
rm.add_link("u2", "g1", "domain1")
84+
rm.add_link("u3", "admin", "domain2")
85+
rm.add_link("u4", "admin", "domain2")
86+
rm.add_link("u4", "admin", "domain1")
87+
rm.add_link("g1", "admin", "domain1")
88+
89+
# Current role inheritance tree:
90+
# domain1:admin domain2:admin
91+
# / \ / \
92+
# domain1:g1 u4 u3
93+
# / \
94+
# u1 u2
95+
96+
self.assertTrue(rm.has_link("u1", "g1", "domain1"))
97+
self.assertFalse(rm.has_link("u1", "g1", "domain2"))
98+
self.assertTrue(rm.has_link("u1", "admin", "domain1"))
99+
self.assertFalse(rm.has_link("u1", "admin", "domain2"))
100+
101+
self.assertTrue(rm.has_link("u2", "g1", "domain1"))
102+
self.assertFalse(rm.has_link("u2", "g1", "domain2"))
103+
self.assertTrue(rm.has_link("u2", "admin", "domain1"))
104+
self.assertFalse(rm.has_link("u2", "admin", "domain2"))
105+
106+
self.assertFalse(rm.has_link("u3", "g1", "domain1"))
107+
self.assertFalse(rm.has_link("u3", "g1", "domain2"))
108+
self.assertFalse(rm.has_link("u3", "admin", "domain1"))
109+
self.assertTrue(rm.has_link("u3", "admin", "domain2"))
110+
111+
self.assertFalse(rm.has_link("u4", "g1", "domain1"))
112+
self.assertFalse(rm.has_link("u4", "g1", "domain2"))
113+
self.assertTrue(rm.has_link("u4", "admin", "domain1"))
114+
self.assertTrue(rm.has_link("u4", "admin", "domain2"))
115+
116+
def test_clear(self):
117+
rm = get_role_manager()
118+
rm.add_link("u1", "g1")
119+
rm.add_link("u2", "g1")
120+
rm.add_link("u3", "g2")
121+
rm.add_link("u4", "g2")
122+
rm.add_link("u4", "g3")
123+
rm.add_link("g1", "g3")
124+
125+
# Current role inheritance tree:
126+
# g3 g2
127+
# / \ / \
128+
# g1 u4 u3
129+
# / \
130+
# u1 u2
131+
132+
rm.clear()
133+
134+
# All data is cleared.
135+
# No role inheritance now.
136+
137+
self.assertFalse(rm.has_link("u1", "g1"))
138+
self.assertFalse(rm.has_link("u1", "g2"))
139+
self.assertFalse(rm.has_link("u1", "g3"))
140+
self.assertFalse(rm.has_link("u2", "g1"))
141+
self.assertFalse(rm.has_link("u2", "g2"))
142+
self.assertFalse(rm.has_link("u2", "g3"))
143+
self.assertFalse(rm.has_link("u3", "g1"))
144+
self.assertFalse(rm.has_link("u3", "g2"))
145+
self.assertFalse(rm.has_link("u3", "g3"))
146+
self.assertFalse(rm.has_link("u4", "g1"))
147+
self.assertFalse(rm.has_link("u4", "g2"))
148+
self.assertFalse(rm.has_link("u4", "g3"))
149+
150+
def test_matching_func(self):
151+
rm = get_role_manager()
152+
rm.add_matching_func(regex_match_func)
153+
154+
rm.add_link("u1", "g1")
155+
rm.add_link("u3", "g2")
156+
rm.add_link("u3", "g3")
157+
rm.add_link(r"u\d+", "g2")
158+
159+
self.assertTrue(rm.has_link("u1", "g1"))
160+
self.assertTrue(rm.has_link("u1", "g2"))
161+
self.assertFalse(rm.has_link("u1", "g3"))
162+
163+
self.assertFalse(rm.has_link("u2", "g1"))
164+
self.assertTrue(rm.has_link("u2", "g2"))
165+
self.assertFalse(rm.has_link("u2", "g3"))
166+
167+
self.assertFalse(rm.has_link("u3", "g1"))
168+
self.assertTrue(rm.has_link("u3", "g2"))
169+
self.assertTrue(rm.has_link("u3", "g3"))
170+
171+
def test_one_to_many(self):
172+
rm = get_role_manager()
173+
rm.add_matching_func(regex_match_func)
174+
175+
rm.add_link("u1", r"g\d+")
176+
self.assertTrue(rm.has_link("u1", "g1"))
177+
self.assertTrue(rm.has_link("u1", "g2"))
178+
self.assertFalse(rm.has_link("u2", "g1"))
179+
self.assertFalse(rm.has_link("u2", "g2"))
180+
181+
def test_many_to_one(self):
182+
rm = get_role_manager()
183+
rm.add_matching_func(regex_match_func)
184+
185+
rm.add_link(r"u\d+", "g1")
186+
self.assertTrue(rm.has_link("u1", "g1"))
187+
self.assertFalse(rm.has_link("u1", "g2"))
188+
self.assertTrue(rm.has_link("u2", "g1"))
189+
self.assertFalse(rm.has_link("u2", "g2"))
190+
191+
def test_matching_func_order(self):
192+
rm = get_role_manager()
193+
rm.add_matching_func(regex_match_func)
194+
195+
rm.add_link(r"g\d+", "root")
196+
rm.add_link("u1", "g1")
197+
self.assertTrue(rm.has_link("u1", "root"))
198+
199+
rm.clear()
200+
201+
rm.add_link("u1", "g1")
202+
rm.add_link(r"g\d+", "root")
203+
self.assertTrue(rm.has_link("u1", "root"))
204+
205+
rm.clear()
206+
207+
rm.add_link("u1", r"g\d+")
208+
rm.add_link("g1", "root")
209+
self.assertTrue(rm.has_link("u1", "root"))
210+
211+
rm.clear()
212+
213+
rm.add_link("g1", "root")
214+
rm.add_link("u1", r"g\d+")
215+
self.assertTrue(rm.has_link("u1", "root"))
216+
217+
def test_concurrent_has_link_with_matching_func(self):
218+
219+
def matching_func(*args):
220+
time.sleep(0.01)
221+
return regex_match_func(*args)
222+
223+
rm = get_role_manager()
224+
rm.add_matching_func(matching_func)
225+
rm.add_link(r"u\d+", "users")
226+
227+
def test_has_link(role):
228+
return rm.has_link(role, "users")
229+
230+
executor = ThreadPoolExecutor(10)
231+
futures = [executor.submit(test_has_link, "u"+str(i)) for i in range(10)]
232+
for future in futures:
233+
self.assertTrue(future.result())
234+

0 commit comments

Comments
 (0)