Support for manual/explicit running of user scripts

Moved user-script endponts under /automation/
Add permission checking for trigger running
This commit is contained in:
Denis Arh
2019-08-23 13:49:36 +02:00
parent 6463df9af1
commit ffdeef1da2
22 changed files with 895 additions and 298 deletions
+22 -4
View File
@@ -1,6 +1,7 @@
package automation
import (
"strconv"
"time"
"github.com/pkg/errors"
@@ -127,12 +128,12 @@ func (s *Script) CheckCompatibility(t *Trigger) error {
return nil
}
// FilterByEvent
// FilterByTrigger
//
// we will use the Trigger struct as a holder for conditions
func (set ScriptSet) FilterByEvent(event, resource string, cc ...TriggerConditionChecker) (out ScriptSet) {
// Filters non-UA scripts that match event and resource + all extra conditions
func (set ScriptSet) FilterByTrigger(event, resource string, cc ...TriggerConditionChecker) (out ScriptSet) {
out, _ = set.Filter(func(s *Script) (bool, error) {
return s.triggers.HasMatch(Trigger{Event: event, Resource: resource}, cc...), nil
return s.IsValid() && s.triggers.HasMatch(Trigger{Event: event, Resource: resource}, cc...), nil
})
return
@@ -170,3 +171,20 @@ func (s *Script) AddTrigger(strategy triggersMergeStrategy, tt ...*Trigger) {
}
}
}
func (s *Script) Triggers() TriggerSet {
return s.triggers
}
func (s Script) HasEvent(event string) bool {
return s.triggers.HasMatch(Trigger{Event: event})
}
func MakeMatcherIDCondition(id uint64) TriggerConditionChecker {
// We'll be comparing strings, not uint64!
var s = strconv.FormatUint(id, 10)
return func(c string) bool {
return s == c
}
}
+91 -66
View File
@@ -33,7 +33,7 @@ type (
}
ScriptsProvider interface {
FilterByEvent(event, resource string, cc ...TriggerConditionChecker) ScriptSet
FilterByTrigger(event, resource string, cc ...TriggerConditionChecker) ScriptSet
}
WatcherService interface {
@@ -66,6 +66,8 @@ func Service(c AutomationServiceConfig) (svc *service) {
trepo: TriggerRepository(c.DbTablePrefix),
db: c.DB,
f: make(chan bool, 64),
}
// Reload ASAP
@@ -74,8 +76,7 @@ func Service(c AutomationServiceConfig) (svc *service) {
}
// Watch watches for changes
func (svc service) Watch(ctx context.Context) {
svc.f = make(chan bool)
func (svc *service) Watch(ctx context.Context) {
go func() {
defer sentry.Recover()
defer close(svc.f)
@@ -102,67 +103,77 @@ func (svc service) Watch(ctx context.Context) {
}
func (svc *service) Reload() {
select {
case svc.f <- true:
return
default:
// that's ok too..
}
go func() {
select {
case svc.f <- true:
return
default:
// that's ok too..
}
}()
}
func (svc *service) reload(ctx context.Context) {
svc.l.Lock()
defer svc.l.Unlock()
if svc.c.DB == nil {
return
}
var (
err error
ss ScriptSet
tt TriggerSet
ss ScriptSet
tt TriggerSet
db = svc.db.With(ctx)
)
ss, err = svc.srepo.findRunnable(svc.db)
svc.logger.Info("scripts loaded", zap.Error(err), zap.Int("count", len(tt)))
if err != nil {
return
}
// Only interested in valid scritps
ss, _ = ss.Filter(func(s *Script) (b bool, e error) {
return s.IsValid(), nil
})
tt, err = svc.trepo.findRunnable(svc.db)
svc.logger.Info("triggers loaded", zap.Error(err), zap.Int("count", len(tt)))
if err != nil {
return
}
_ = tt.Walk(func(t *Trigger) error {
s := ss.FindByID(t.ScriptID)
if s != nil && t.IsValid() && s.CheckCompatibility(t) != nil {
// Add only compatible triggers
s.triggers = append(s.triggers, t)
_ = db.Transaction(func() (err error) {
ss, err = svc.srepo.findRunnable(db)
svc.logger.Info("scripts loaded", zap.Error(err), zap.Int("count", len(ss)))
if err != nil {
return
}
return nil
// Remove all invalid scripts
svc.runnables, _ = ss.Filter(func(s *Script) (b bool, e error) {
return s.IsValid(), nil
})
tt, err = svc.trepo.findRunnable(db)
svc.logger.Info("triggers loaded", zap.Error(err), zap.Int("count", len(tt)))
if err != nil {
return
}
return tt.Walk(func(t *Trigger) error {
s := svc.runnables.FindByID(t.ScriptID)
if s != nil && t.IsValid() && s.CheckCompatibility(t) == nil {
// Add only compatible triggers
s.triggers = append(s.triggers, t)
}
return nil
})
})
}
// FindRunnableScripts scans internal list of runnable scripts and filters them by (trigger's) event and origin
func (svc service) FindRunnableScripts(event, origin string, cc ...TriggerConditionChecker) ScriptSet {
return svc.runnables.FilterByEvent(event, origin, cc...)
// FindRunnableScripts finds runnable scripts in internal list
//
// It uses resource, event and extra condition checkers to filter out all scripts
// that have matching triggers
func (svc service) FindRunnableScripts(resource, event string, cc ...TriggerConditionChecker) ScriptSet {
svc.l.Lock()
defer svc.l.Unlock()
return svc.runnables.FilterByTrigger(
event,
resource,
cc...,
)
}
func (svc service) FindScriptByID(ctx context.Context, scriptID uint64) (*Script, error) {
return svc.srepo.findByID(svc.db, scriptID)
return svc.srepo.findByID(svc.db.With(ctx), scriptID)
}
func (svc service) FindScripts(ctx context.Context, f ScriptFilter) (ScriptSet, ScriptFilter, error) {
return svc.srepo.find(svc.db, f)
return svc.srepo.find(svc.db.With(ctx), f)
}
// CreateScript - modifies script's props, pushes to repo & updates scripts cache
@@ -171,8 +182,13 @@ func (svc service) CreateScript(ctx context.Context, s *Script) error {
s.CreatedAt = time.Now()
s.CreatedBy = auth.GetIdentityFromContext(ctx).Identity()
return svc.db.Transaction(func() (err error) {
if err = svc.srepo.create(svc.db, s); err != nil {
db := svc.db.With(ctx)
// Reloading scripts at the end (after the transaction completes)
defer svc.Reload()
return db.Transaction(func() (err error) {
if err = svc.srepo.create(db, s); err != nil {
return
}
@@ -185,11 +201,10 @@ func (svc service) CreateScript(ctx context.Context, s *Script) error {
}
// Force no-pre-check
if err = svc.trepo.mergeSet(svc.db, STMS_FRESH, s.ID, s.triggers); err != nil {
if err = svc.trepo.mergeSet(db, STMS_FRESH, s.ID, s.triggers); err != nil {
return
}
svc.Reload()
return
})
}
@@ -201,8 +216,13 @@ func (svc service) UpdateScript(ctx context.Context, s *Script) error {
*s.UpdatedAt = time.Now()
s.DeletedAt, s.DeletedBy = nil, 0
return svc.db.Transaction(func() (err error) {
if err = svc.srepo.update(svc.db, s); err != nil {
db := svc.db.With(ctx)
// Reloading scripts at the end (after the transaction completes)
defer svc.Reload()
return db.Transaction(func() (err error) {
if err = svc.srepo.update(db, s); err != nil {
return
}
@@ -218,11 +238,10 @@ func (svc service) UpdateScript(ctx context.Context, s *Script) error {
return
}
if err = svc.trepo.mergeSet(svc.db, s.tms, s.ID, s.triggers); err != nil {
if err = svc.trepo.mergeSet(db, s.tms, s.ID, s.triggers); err != nil {
return
}
svc.Reload()
return
})
}
@@ -233,25 +252,31 @@ func (svc service) DeleteScript(ctx context.Context, s *Script) (err error) {
*s.DeletedAt = time.Now()
s.DeletedBy = auth.GetIdentityFromContext(ctx).Identity()
// We're doing soft delete in the repo
if err = svc.srepo.update(svc.db, s); err != nil {
return err
}
db := svc.db.With(ctx)
if err = svc.trepo.deleteByScriptID(svc.db, s.ID); err != nil {
return err
}
// Reloading scripts at the end (after the transaction completes)
defer svc.Reload()
svc.Reload()
return
return db.Transaction(func() error {
// We're doing soft delete in the repo
if err = svc.srepo.update(db, s); err != nil {
return err
}
if err = svc.trepo.deleteByScriptID(db, s.ID); err != nil {
return err
}
return nil
})
}
func (svc service) FindTriggerByID(ctx context.Context, scriptID uint64) (*Trigger, error) {
return svc.trepo.findByID(svc.db, scriptID)
return svc.trepo.findByID(svc.db.With(ctx), scriptID)
}
func (svc service) FindTriggers(ctx context.Context, f TriggerFilter) (TriggerSet, TriggerFilter, error) {
return svc.trepo.find(svc.db, f)
return svc.trepo.find(svc.db.With(ctx), f)
}
// CreateTrigger - modifies script's props, pushes to repo & updates scripts cache
@@ -260,7 +285,7 @@ func (svc service) CreateTrigger(ctx context.Context, s *Script, t *Trigger) (er
return err
}
if err = svc.trepo.replace(svc.db, t); err != nil {
if err = svc.trepo.replace(svc.db.With(ctx), t); err != nil {
return err
}
@@ -286,7 +311,7 @@ func (svc service) UpdateTrigger(ctx context.Context, s *Script, t *Trigger) (er
return err
}
if err = svc.trepo.replace(svc.db, t); err != nil {
if err = svc.trepo.replace(svc.db.With(ctx), t); err != nil {
return err
}
@@ -316,7 +341,7 @@ func (svc service) DeleteTrigger(ctx context.Context, t *Trigger) (err error) {
t.DeletedBy = auth.GetIdentityFromContext(ctx).Identity()
// We're doing soft delete in the repo
if err = svc.trepo.replace(svc.db, t); err != nil {
if err = svc.trepo.replace(svc.db.With(ctx), t); err != nil {
return err
}
+2 -2
View File
@@ -95,12 +95,12 @@ withTriggers:
continue withTriggers
}
if m.Resource != t.Resource {
if m.Resource != "" && m.Resource != t.Resource {
// event should match
continue withTriggers
}
if m.Event != t.Event {
if m.Event != "" && m.Event != t.Event {
// event should match
continue withTriggers
}
+1 -1
View File
@@ -22,7 +22,7 @@ func TestTriggerSet_HasMatch(t *testing.T) {
}, {
name: "simple miss",
set: TriggerSet{nil, &Trigger{}, &Trigger{Event: "e", Enabled: true}, nil, &Trigger{}},
args: args{m: Trigger{}},
args: args{m: Trigger{Event: "E"}},
want: false,
}, {
name: "specific",