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:
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user