From 2cc865abd66c49ea2c200740f6550dcc932954d9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Toma=C5=BE=20Jerman?= Date: Sat, 30 Nov 2024 11:50:26 +0100 Subject: [PATCH] Tweak svc, add tests --- server/pkg/rbac/rule_index.go | 4 + server/pkg/rbac/service.go | 67 ++++++---- server/pkg/rbac/wrapper_test.go | 213 ++++++++++++++++++++++++++++++++ 3 files changed, 263 insertions(+), 21 deletions(-) create mode 100644 server/pkg/rbac/wrapper_test.go diff --git a/server/pkg/rbac/rule_index.go b/server/pkg/rbac/rule_index.go index f0f1678d8..e2f00fd7c 100644 --- a/server/pkg/rbac/rule_index.go +++ b/server/pkg/rbac/rule_index.go @@ -99,6 +99,10 @@ func (t *ruleIndex) collect(exact bool, role uint64, op, res string) (out []*Rul // An edge case implied by the test suite if op == "" && res == "" { + if t.children[role].children[""] == nil || t.children[role].children[""].children[""] == nil { + return + } + out = append(out, t.children[role].children[""].children[""].rule) return } diff --git a/server/pkg/rbac/service.go b/server/pkg/rbac/service.go index a4dc96ff8..95944a80f 100644 --- a/server/pkg/rbac/service.go +++ b/server/pkg/rbac/service.go @@ -178,9 +178,28 @@ func NoopSvc(access Access) (svc *Service) { // NewService initializes the wrapper service with all the required surrounding bits func NewService(ctx context.Context, l *zap.Logger, store rbacRulesStore, cc Config) (svc *Service, err error) { - cc = defaultWrapperConfig(l, cc) + cc = defaultWrapperConfig(cc) - usageCounter := &usageCounter[string]{ + uc := initUsageCounter(ctx, cc) + sl := initStatsLogger(ctx, l) + svc = initSvc(ctx, l, cc, sl, uc) + + // Init bits and pieces + svc.roles, err = svc.loadRoles(ctx) + if err != nil { + return + } + + svc.index, err = svc.loadIndex(ctx) + if err != nil { + return + } + + return +} + +func initUsageCounter(ctx context.Context, cc Config) (svc *usageCounter[string]) { + svc = &usageCounter[string]{ incChan: make(chan string, 1024), decayFactor: cc.DecayFactor, @@ -192,42 +211,41 @@ func NewService(ctx context.Context, l *zap.Logger, store rbacRulesStore, cc Con }, } - sl := &statsLogger{ + svc.watch(ctx) + return +} + +func initStatsLogger(ctx context.Context, l *zap.Logger) (svc *statsLogger) { + svc = &statsLogger{ log: l.Named("rbac stats logger"), cacheHitChan: make(chan statsWrap, 1024), cacheMissChan: make(chan statsWrap, 1024), timingChan: make(chan time.Duration, 1024), } + svc.watch(ctx) + return +} + +func initSvc(ctx context.Context, l *zap.Logger, cc Config, sl *statsLogger, uc *usageCounter[string]) (svc *Service) { svc = &Service{ + logger: l, + cfg: cc, StatLogger: sl, - logger: l, - usageCounter: usageCounter, + usageCounter: uc, RuleStorage: cc.RuleStorage, RoleStorage: cc.RoleStorage, } - svc.roles, err = svc.loadRoles(ctx) - if err != nil { - return - } - - svc.index, err = svc.loadIndex(ctx, store, svc.roles) - if err != nil { - return - } - - usageCounter.watch(ctx) svc.watch(ctx) - sl.watch(ctx) return } -func defaultWrapperConfig(l *zap.Logger, base Config) (out Config) { +func defaultWrapperConfig(base Config) (out Config) { out = base // -1 disables partitioning so everything is pulled in memory @@ -830,9 +848,16 @@ func (svc *Service) getMatchingRule(st evaluationState, kind roleKind, role uint // segmentRoles determines what roles are indexed and unindexed func (svc *Service) segmentRoles(roles partRoles, resource string) (indexed, unindexed partRoles, err error) { + svc.mux.RLock() + defer svc.mux.RUnlock() + unindexed = partRoles{} indexed = partRoles{} + if svc.index.index.empty() { + return indexed, roles, nil + } + unindexed[CommonRole] = make(map[uint64]bool) indexed[CommonRole] = make(map[uint64]bool) @@ -958,7 +983,7 @@ func (svc *Service) updateWrapperIndexMemFirst(ctx context.Context) (err error) return } - svc.swapIndexes(ctx, auxIndex) + svc.swapIndexes(auxIndex) return } @@ -1000,7 +1025,7 @@ func (svc *Service) indexForResources(ctx context.Context, res ...string) (index return } -func (svc *Service) loadIndex(ctx context.Context, s rbacRulesStore, allRoles []*Role) (out *wrapperIndex, err error) { +func (svc *Service) loadIndex(ctx context.Context) (out *wrapperIndex, err error) { // How do we figure out what resources we have? // do we just start from empty? @@ -1029,7 +1054,7 @@ func (svc *Service) buildNewIndex(ctx context.Context) (index *wrapperIndex, err return svc.indexForResources(ctx, res...) } -func (svc *Service) swapIndexes(ctx context.Context, auxIndex *wrapperIndex) { +func (svc *Service) swapIndexes(auxIndex *wrapperIndex) { if auxIndex == nil { return } diff --git a/server/pkg/rbac/wrapper_test.go b/server/pkg/rbac/wrapper_test.go new file mode 100644 index 000000000..68ac17f29 --- /dev/null +++ b/server/pkg/rbac/wrapper_test.go @@ -0,0 +1,213 @@ +package rbac + +import ( + "context" + "testing" + + "github.com/cortezaproject/corteza/server/system/types" + "github.com/stretchr/testify/require" +) + +func TestRoleSegmentation(t *testing.T) { + req := require.New(t) + + wx := &wrapperIndex{} + w := Service{ + index: wx, + } + + rl1 := uint64(1001) + rl2 := uint64(2001) + res1 := "abc/1/2/3" + res2 := "def/1/2/3" + + wx.add(rl1, res1, &Rule{ + RoleID: rl1, + Resource: res1, + Operation: "read", + Access: Allow, + }) + + rls := partRoles{} + rls[CommonRole] = map[uint64]bool{ + rl1: true, + rl2: true, + } + + indexed, unindexed, err := w.segmentRoles(rls, res1) + req.NoError(err) + + req.True(indexed[CommonRole][rl1]) + req.False(indexed[CommonRole][rl2]) + + req.True(unindexed[CommonRole][rl2]) + req.False(unindexed[CommonRole][rl1]) + + // + // + + indexed, unindexed, err = w.segmentRoles(rls, res2) + req.NoError(err) + + req.False(indexed[CommonRole][rl1]) + req.False(indexed[CommonRole][rl2]) + + req.True(unindexed[CommonRole][rl1]) + req.True(unindexed[CommonRole][rl2]) +} + +func TestRoleSegmentationEmpty(t *testing.T) { + req := require.New(t) + + wx := &wrapperIndex{} + w := Service{ + index: wx, + } + + rl1 := uint64(1001) + rl2 := uint64(2001) + res1 := "abc/1/2/3" + + rls := partRoles{} + rls[CommonRole] = map[uint64]bool{ + rl1: true, + rl2: true, + } + + _, unindexed, err := w.segmentRoles(rls, res1) + req.NoError(err) + + req.True(unindexed[CommonRole][rl1]) + req.True(unindexed[CommonRole][rl2]) +} + +type ( + tRuleStore struct { + searches []RuleFilter + } +) + +func TestPullRules(t *testing.T) { + req := require.New(t) + ruleS := &tRuleStore{} + ctx := context.Background() + + wx := &Service{ + RuleStorage: ruleS, + } + + wx.pullRules(ctx, 1, "res/1/2/3") + req.Len(ruleS.searches, 1) + req.Equal([]string{"res/1/2/3", "res/1/2/*", "res/1/*/*", "res/*/*/*"}, ruleS.searches[0].Resource) + req.Equal(uint64(1), ruleS.searches[0].RoleID) + + wx.pullRules(ctx, 1, "res/1") + req.Len(ruleS.searches, 2) + req.Equal([]string{"res/1", "res/*"}, ruleS.searches[1].Resource) + req.Equal(uint64(1), ruleS.searches[1].RoleID) + + wx.pullRules(ctx, 1, "res") + req.Len(ruleS.searches, 3) + req.Equal([]string{"res"}, ruleS.searches[2].Resource) + req.Equal(uint64(1), ruleS.searches[2].RoleID) +} + +func TestCombiningSources(t *testing.T) { + req := require.New(t) + wx := &Service{ + index: &wrapperIndex{}, + } + + wx.index.add(1, "res/1/2/3", &Rule{ + RoleID: 1, + Resource: "res/1/2/3", + Operation: "read", + Access: Inherit, + }, &Rule{ + RoleID: 1, + Resource: "res/1/2/*", + Operation: "read", + Access: Allow, + }, &Rule{ + RoleID: 2, + Resource: "res/1/2/3", + Operation: "read", + Access: Deny, + }) + + stt := evaluationState{ + res: "res/1/2/3", + op: "read", + + unindexedRoles: partRoles{CommonRole: map[uint64]bool{3: true}}, + indexedRoles: partRoles{CommonRole: map[uint64]bool{1: true}}, + unindexedRules: [5]map[uint64][]*Rule{CommonRole: { + 3: {{ + RoleID: 3, + Resource: "res/1/2/3", + Operation: "read", + Access: Inherit, + }, { + RoleID: 3, + Resource: "res/1/2/*", + Operation: "read", + Access: Deny, + }}, + }}, + } + + auxRule := wx.getMatchingRule(stt, CommonRole, 1) + req.Equal("res/1/2/*", auxRule.Resource) + req.Equal(Allow, auxRule.Access) + + auxRule = wx.getMatchingRule(stt, CommonRole, 3) + req.Equal("res/1/2/*", auxRule.Resource) + req.Equal(Deny, auxRule.Access) + + wx.index.add(3, "res/1/2/3", &Rule{ + RoleID: 3, + Resource: "res/1/2/3", + Operation: "read", + Access: Inherit, + }) + stt = evaluationState{ + res: "res/1/2/3", + op: "read", + + unindexedRoles: partRoles{CommonRole: map[uint64]bool{3: true}}, + indexedRoles: partRoles{CommonRole: map[uint64]bool{1: true}}, + unindexedRules: [5]map[uint64][]*Rule{CommonRole: { + 3: {{ + RoleID: 3, + Resource: "res/1/2/*", + Operation: "read", + Access: Deny, + }}, + }}, + } + + auxRule = wx.getMatchingRule(stt, CommonRole, 3) + req.Equal("res/1/2/*", auxRule.Resource) + req.Equal(Deny, auxRule.Access) +} + +func (tt *tRuleStore) SearchRbacRules(ctx context.Context, f RuleFilter) (rs RuleSet, rf RuleFilter, err error) { + tt.searches = append(tt.searches, f) + return +} + +func (tt *tRuleStore) UpsertRbacRule(ctx context.Context, rr ...*Rule) (err error) { + return +} + +func (tt *tRuleStore) DeleteRbacRule(ctx context.Context, rr ...*Rule) (err error) { + return +} + +func (tt *tRuleStore) TruncateRbacRules(ctx context.Context) (err error) { + return +} + +func (tt *tRuleStore) SearchRoles(ctx context.Context, f types.RoleFilter) (rs types.RoleSet, rf types.RoleFilter, err error) { + return +}