diff --git a/internal/permissions/permissions.go b/internal/permissions/permissions.go index 4e0f32073..c4efa0e9e 100644 --- a/internal/permissions/permissions.go +++ b/internal/permissions/permissions.go @@ -10,7 +10,13 @@ type ( CheckAccessFunc func() Access ) -const EveryoneRoleID = 1 +const ( + // Hardcoded Role ID for everyone + EveryoneRoleID = 1 + + // Hardcoded ID for Admin role + AdminRoleID = 2 +) func (a Access) String() string { switch a { diff --git a/internal/permissions/repository.go b/internal/permissions/repository.go index fabbbca6a..b2dac2f2b 100644 --- a/internal/permissions/repository.go +++ b/internal/permissions/repository.go @@ -3,7 +3,9 @@ package permissions import ( "context" + "github.com/pkg/errors" "github.com/titpetric/factory" + "gopkg.in/Masterminds/squirrel.v1" ) type ( @@ -32,7 +34,7 @@ func (r repository) columns() []string { "rel_role", "resource", "operation", - "value", + "access", } } @@ -43,7 +45,40 @@ func (r *repository) With(ctx context.Context) *repository { } } -func (r *repository) Load() (rr RuleSet, err error) { - // @todo load and return - return nil, nil +func (r *repository) Load() (RuleSet, error) { + rr := make([]*Rule, 0) + + lookup := squirrel. + Select(r.columns()...). + From(r.dbTable) + + if query, args, err := lookup.ToSql(); err != nil { + return nil, errors.Wrap(err, "could not build lookup query for permission rules") + } else if err = r.dbh.Select(&rr, query, args...); err != nil { + return nil, errors.Wrap(err, "could not get permission rules") + } + + return rr, nil +} + +func (r *repository) Store(deleteSet, updateSet RuleSet) (err error) { + return r.dbh.Transaction(func() error { + if len(deleteSet) > 0 { + err = r.dbh.Delete(r.dbTable, deleteSet, "rel_role", "resource", "operation") + if err != nil { + return err + } + } + + if len(updateSet) > 0 { + err = updateSet.Walk(func(rule *Rule) error { + return r.dbh.Replace(r.dbTable, rule) + }) + if err != nil { + return err + } + } + + return nil + }) } diff --git a/internal/permissions/rule.go b/internal/permissions/rule.go index 7aa8623ec..e1f6e7105 100644 --- a/internal/permissions/rule.go +++ b/internal/permissions/rule.go @@ -9,7 +9,7 @@ type ( RoleID uint64 `json:"roleID,string" db:"rel_role"` Resource Resource `json:"resource" db:"resource"` Operation Operation `json:"operation" db:"operation"` - Access Access `json:"value,string" db:"value"` + Access Access `json:"access,string" db:"access"` } ) @@ -29,6 +29,10 @@ func (r Rule) String() string { } func (r Rule) Equals(cmp *Rule) bool { + if cmp == nil { + return false + } + return r.RoleID == cmp.RoleID && r.Resource == cmp.Resource && r.Operation == cmp.Operation diff --git a/internal/permissions/ruleset_utils.go b/internal/permissions/ruleset_utils.go new file mode 100644 index 000000000..8c095e61a --- /dev/null +++ b/internal/permissions/ruleset_utils.go @@ -0,0 +1,54 @@ +package permissions + +func (set RuleSet) merge(rules ...*Rule) (out RuleSet, err error) { + var ( + o int + olen = len(set) + + skipInherited = func(r *Rule) (b bool, e error) { + return r != nil, nil + } + + merged = set + ) + + if olen == 0 { + // Nothing exists yet, just assign + merged = rules + } else { + newRules: + for _, rule := range rules { + for ; o < olen; o++ { + // Never go beyond the last old rule + if merged[o].Equals(rule) { + merged[o].Access = rule.Access + + // only one rule can match so proceed with next new rule + continue newRules + } + } + + // none of the old rules matched, append + merged = append(merged, rule) + } + + } + + // Filter out all rules with access = inherit + return merged.Filter(skipInherited) +} + +func (set RuleSet) split() (inherited, rest RuleSet) { + inherited, rest = RuleSet{}, RuleSet{} + + for _, r := range set { + if r.Access == Inherit { + inherited = append(inherited, r) + } else { + rest = append(rest, r) + + } + } + + return +} diff --git a/internal/permissions/ruleset_utils_test.go b/internal/permissions/ruleset_utils_test.go new file mode 100644 index 000000000..f1a047453 --- /dev/null +++ b/internal/permissions/ruleset_utils_test.go @@ -0,0 +1,90 @@ +package permissions + +import ( + "reflect" + "testing" + + "github.com/crusttech/crust/internal/test" +) + +// Test role inheritance +func TestRuleSet_merge(t *testing.T) { + var ( + assert = test.Assert + + sCases = []struct { + old RuleSet + in RuleSet + exp RuleSet + }{ + { + RuleSet{ + &Rule{role1, resService1, opAccess, Allow}, + &Rule{role2, resService1, opAccess, Deny}, + &Rule{EveryoneRoleID, resService2, opAccess, Deny}, + &Rule{role1, resService2, opAccess, Allow}, + }, + RuleSet{ + &Rule{EveryoneRoleID, resThingWc, opAccess, Deny}, + &Rule{role1, resThing42, opAccess, Allow}, + &Rule{role1, resThing42, opAccess, Inherit}, + }, + RuleSet{ + &Rule{role1, resService1, opAccess, Allow}, + &Rule{role2, resService1, opAccess, Deny}, + &Rule{EveryoneRoleID, resService2, opAccess, Deny}, + &Rule{role1, resService2, opAccess, Allow}, + &Rule{EveryoneRoleID, resThingWc, opAccess, Deny}, + &Rule{role1, resThing42, opAccess, Allow}, + &Rule{role1, resThing42, opAccess, Inherit}, + }, + }, + } + ) + + for c, sc := range sCases { + out, _ := sc.old.merge(sc.in...) + + assert(t, len(out) == len(sc.exp), "Check test #%d failed, expected length %d, got %d", c, len(out), len(sc.exp)) + assert(t, reflect.DeepEqual(out, sc.exp), "Check test #%d failed, reflect.DeepEqual == false", c) + + } +} + +// Test role inheritance +func TestRuleSet_split(t *testing.T) { + var ( + assert = test.Assert + + sCases = []struct { + set RuleSet + i RuleSet + r RuleSet + }{ + { + RuleSet{ + &Rule{role1, resService1, opAccess, Allow}, + &Rule{role2, resService1, opAccess, Deny}, + &Rule{EveryoneRoleID, resService2, opAccess, Inherit}, + }, + RuleSet{ + &Rule{EveryoneRoleID, resService2, opAccess, Inherit}, + }, + RuleSet{ + &Rule{role1, resService1, opAccess, Allow}, + &Rule{role2, resService1, opAccess, Deny}, + }, + }, + } + ) + + for c, sc := range sCases { + i, r := sc.set.split() + + assert(t, len(i) == len(sc.i), "Check test #%d failed, expected length %d, got %d", c, len(i), len(sc.i)) + assert(t, len(r) == len(sc.r), "Check test #%d failed, expected length %d, got %d", c, len(r), len(sc.r)) + assert(t, reflect.DeepEqual(i, sc.i), "Check test #%d failed, reflect.DeepEqual == false", c) + assert(t, reflect.DeepEqual(r, sc.r), "Check test #%d failed, reflect.DeepEqual == false", c) + + } +} diff --git a/internal/permissions/service.go b/internal/permissions/service.go index ccf0e8663..b39c41b60 100644 --- a/internal/permissions/service.go +++ b/internal/permissions/service.go @@ -3,41 +3,42 @@ package permissions import ( "context" "sync" + "time" + + "go.uber.org/zap" ) type ( service struct { - l sync.Locker + l sync.Mutex + logger *zap.Logger + + // service will flush values on TRUE or just reload on FALSE + f chan bool rules RuleSet repository *repository } +) - Verifier interface { - Can(ctx context.Context, res Resource, op Operation, ff ...CheckAccessFunc) bool - } +const ( + watchInterval = time.Second * 60 ) // Service initializes service{} struct // // service{} struct preloads, checks, grants and flushes privileges to and from repository // It acts as a caching layer -func Service(repository *repository) *service { - return &service{ +func Service(ctx context.Context, logger *zap.Logger, repository *repository) (svc *service) { + svc = &service{ + f: make(chan bool, 0), + + logger: logger.Named("permissions"), repository: repository, } -} -func (svc *service) Preload(ctx context.Context) (err error) { - svc.l.Lock() - defer svc.l.Unlock() - - svc.rules, err = svc.repository.With(ctx).Load() - if err != nil { - return - } - - return nil + svc.Reload(ctx) + return } // Can function performs permission check for roles in context @@ -82,15 +83,67 @@ func (svc service) Check(res Resource, op Operation, roles ...uint64) (v Access) // Grant appends and/or overwrites internal rules slice // // All rules with Inherit are removed -func (svc service) Grant(ctx context.Context, rules ...*Rule) error { +func (svc *service) Grant(ctx context.Context, rules ...*Rule) (err error) { svc.l.Lock() defer svc.l.Unlock() - // @todo update svc.rules + if svc.rules, err = svc.rules.merge(rules...); err != nil { + return + } - return nil + return svc.flush(ctx) } -func (svc service) watcher() { - // @todo will listen to chan and load new stuff every time it gets a ping +// Watches for changes +func (svc service) Watch(ctx context.Context) { + go func() { + var ticker = time.NewTicker(watchInterval) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + break + case <-ticker.C: + svc.Reload(ctx) + case <-svc.f: + svc.Reload(ctx) + } + } + }() + + svc.logger.Debug("watcher initialized") +} + +func (svc service) Reload(ctx context.Context) { + svc.l.Lock() + defer svc.l.Unlock() + + rr, err := svc.repository.With(ctx).Load() + svc.logger.Info( + "reloading rules", + zap.Error(err), + zap.Int("before", len(svc.rules)), + zap.Int("after", len(rr)), + ) + + if err != nil { + svc.rules = rr + } +} + +func (svc service) flush(ctx context.Context) (err error) { + d, u := svc.rules.split() + err = svc.repository.With(ctx).Store(d, u) + + if err != nil { + return + } + + svc.rules = u + svc.logger.Info("flushed rules", + zap.Int("updated", len(u)), + zap.Int("deleted", len(d))) + + return }