From 5afba0eb196fb8f71bddac964789d91cc21332e0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Toma=C5=BE=20Jerman?= Date: Tue, 13 Aug 2024 15:00:10 +0200 Subject: [PATCH] update rule index to allow modifications --- server/pkg/rbac/rule_index.go | 91 +++++++++++++++-- server/pkg/rbac/rule_index_test.go | 156 ++++++++++++++++++++++++++++- 2 files changed, 233 insertions(+), 14 deletions(-) diff --git a/server/pkg/rbac/rule_index.go b/server/pkg/rbac/rule_index.go index 2198f174d..07b8aae4d 100644 --- a/server/pkg/rbac/rule_index.go +++ b/server/pkg/rbac/rule_index.go @@ -20,6 +20,8 @@ type ( isLeaf bool access Access rule *Rule + + count int } ) @@ -27,16 +29,28 @@ type ( // // The build isn't that cleanned up but the lookup is good, I promise <3 func buildRuleIndex(rules []*Rule) (index *ruleIndex) { - index = &ruleIndex{ - children: make(map[uint64]*ruleIndexNode, 8), + index = &ruleIndex{} + index.add(rules...) + return index +} + +func (index *ruleIndex) add(rules ...*Rule) { + if index.children == nil { + index.children = make(map[uint64]*ruleIndexNode, len(rules)/2) } for _, r := range rules { + // skip duplicates + if index.has(r) { + continue + } + if _, ok := index.children[r.RoleID]; !ok { index.children[r.RoleID] = &ruleIndexNode{ children: make(map[string]*ruleIndexNode, 4), } } + index.children[r.RoleID].count++ n := index.children[r.RoleID] bits := append([]string{r.Operation}, strings.Split(r.Resource, "/")...) @@ -46,6 +60,7 @@ func buildRuleIndex(rules []*Rule) (index *ruleIndex) { children: make(map[string]*ruleIndexNode, 4), } } + n.children[b].count++ n = n.children[b] } @@ -54,8 +69,66 @@ func buildRuleIndex(rules []*Rule) (index *ruleIndex) { n.access = r.Access n.rule = r } +} - return index +func (index *ruleIndex) remove(rules ...*Rule) { + if len(rules) == 0 { + return + } + + for _, r := range rules { + if _, ok := index.children[r.RoleID]; !ok { + continue + } + + if !index.has(r) { + continue + } + + bits := append([]string{r.Operation}, strings.Split(r.Resource, "/")...) + index.removeRec(index.children[r.RoleID], bits) + + // Finishing touch cleanup + if len(index.children[r.RoleID].children[r.Operation].children) == 0 { + delete(index.children[r.RoleID].children, r.Operation) + } + if len(index.children[r.RoleID].children) == 0 { + delete(index.children, r.RoleID) + } + } +} + +func (index *ruleIndex) removeRec(n *ruleIndexNode, bits []string) { + // Recursive in; decrement counters + n.count-- + + if len(bits) == 0 { + return + } + + n = n.children[bits[0]] + index.removeRec(n, bits[1:]) + + // Recursive out; yoink out obsolete štuff + if len(bits) == 1 { + if len(n.children) > 0 { + n.children[bits[0]].isLeaf = false + n.children[bits[0]].rule = nil + } + return + } + + if n.children[bits[1]].count == 0 { + delete(n.children, bits[1]) + } +} + +func (t *ruleIndex) has(r *Rule) bool { + return len(t.collect(true, r.RoleID, r.Operation, r.Resource)) > 0 +} + +func (t *ruleIndex) get(role uint64, op, res string) (out []*Rule) { + return t.collect(false, role, op, res) } // get returns all RBAC rules matching these constraints @@ -64,7 +137,7 @@ func buildRuleIndex(rules []*Rule) (index *ruleIndex) { // the operation + 1 for the role. // // Our longest bit will be 6 so this is essentially constant time. -func (t *ruleIndex) get(role uint64, op, res string) (out []*Rule) { +func (t *ruleIndex) collect(exact bool, role uint64, op, res string) (out []*Rule) { if t.children == nil { return } @@ -90,7 +163,7 @@ func (t *ruleIndex) get(role uint64, op, res string) (out []*Rule) { return } - return aux.get(res, 0) + return aux.get(exact, res, 0) } // get returns all of the rules matching these constraints @@ -100,7 +173,7 @@ func (t *ruleIndex) get(role uint64, op, res string) (out []*Rule) { // be a memory hog in scenarios where we're pounding this function. // // The from denotes the substring we've not yet processed. -func (n *ruleIndexNode) get(res string, from int) (out []*Rule) { +func (n *ruleIndexNode) get(exact bool, res string, from int) (out []*Rule) { if n == nil || n.children == nil { return } @@ -133,12 +206,12 @@ func (n *ruleIndexNode) get(res string, from int) (out []*Rule) { // Get RBAC rules down the actual path pathBit := res[from:nextDelim] if n.children[pathBit] != nil { - out = append(out, n.children[pathBit].get(res, nextDelim+1)...) + out = append(out, n.children[pathBit].get(exact, res, nextDelim+1)...) } // Get RBAC rules down the wildcard path - if n.children[wildcard] != nil { - out = append(out, n.children[wildcard].get(res, nextDelim+1)...) + if !exact && n.children[wildcard] != nil { + out = append(out, n.children[wildcard].get(exact, res, nextDelim+1)...) } return diff --git a/server/pkg/rbac/rule_index_test.go b/server/pkg/rbac/rule_index_test.go index 3390673a9..e063b737f 100644 --- a/server/pkg/rbac/rule_index_test.go +++ b/server/pkg/rbac/rule_index_test.go @@ -7,11 +7,13 @@ import ( "github.com/stretchr/testify/require" ) -func TestIndex(t *testing.T) { +func TestIndexBuild(t *testing.T) { tcc := []struct { - name string - in []*Rule - out []int + name string + in []*Rule + remove []*Rule + add []*Rule + out []int role uint64 op string @@ -118,15 +120,159 @@ func TestIndex(t *testing.T) { role: 1, op: "read", res: "a:b/c/d", + }, { + name: "removing the only element", + in: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }}, + remove: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }}, + out: nil, + + role: 1, + op: "write", + res: "a:b/c/d", + }, + { + name: "removing twice added thing", + in: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }, { + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }}, + remove: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }}, + + out: nil, + + role: 1, + op: "write", + res: "a:b/c/d", + }, + { + name: "two elements with no common root", + in: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }, { + RoleID: 2, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }}, + remove: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }}, + out: nil, + + role: 1, + op: "write", + res: "a:b/c/d", + }, + { + name: "two elements with common root (get removed)", + in: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }, { + RoleID: 1, + Resource: "a:b/c/e", + Operation: "write", + Access: Allow, + }}, + remove: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }}, + out: nil, + + role: 1, + op: "write", + res: "a:b/c/d", + }, + { + name: "two elements with common root (get not removed)", + in: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }, { + RoleID: 1, + Resource: "a:b/c/e", + Operation: "write", + Access: Allow, + }}, + remove: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }}, + out: []int{1}, + + role: 1, + op: "write", + res: "a:b/c/e", + }, + { + name: "add new element", + in: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/d", + Operation: "write", + Access: Allow, + }}, + add: []*Rule{{ + RoleID: 1, + Resource: "a:b/c/x", + Operation: "write", + Access: Allow, + }}, + + out: []int{1}, + + role: 1, + op: "write", + res: "a:b/c/x", }} for _, tc := range tcc { t.Run(tc.name, func(t *testing.T) { ix := buildRuleIndex(tc.in) + ix.remove(tc.remove...) + ix.add(tc.add...) + out := RuleSet(ix.get(tc.role, tc.op, tc.res)) sort.Sort(out) - want := RuleSet(graby(tc.in, tc.out)) + want := RuleSet(graby(append(tc.in, tc.add...), tc.out)) sort.Sort(want) require.Len(t, out, len(want))