diff --git a/.github/workflows/codetests.yml b/.github/workflows/codetests.yml index 05e8729..ee21506 100644 --- a/.github/workflows/codetests.yml +++ b/.github/workflows/codetests.yml @@ -13,10 +13,10 @@ jobs: os: [macos, windows, ubuntu] runs-on: ${{ matrix.os }}-latest steps: - - uses: actions/checkout@v4 - - uses: actions/setup-go@v4 + - uses: actions/checkout@v6 + - uses: actions/setup-go@v6 with: - go-version: 1.19 + go-version: stable - name: go-test run: go test -race -covermode=atomic '-test.v' ./... # Runs golangci-lint on macos against freebsd and macos. @@ -29,14 +29,14 @@ jobs: env: GOOS: ${{ matrix.os }} steps: - - uses: actions/setup-go@v4 + - uses: actions/setup-go@v6 with: - go-version: 1.19 - - uses: actions/checkout@v4 + go-version: stable + - uses: actions/checkout@v6 - name: golangci-lint - uses: golangci/golangci-lint-action@v3 + uses: golangci/golangci-lint-action@v9 with: - version: v1.50 + version: v2.9 # Runs golangci-lint on linux against linux and windows. golangci-linux: strategy: @@ -47,11 +47,11 @@ jobs: env: GOOS: ${{ matrix.os }} steps: - - uses: actions/setup-go@v4 + - uses: actions/setup-go@v6 with: - go-version: 1.19 - - uses: actions/checkout@v4 + go-version: stable + - uses: actions/checkout@v6 - name: golangci-lint - uses: golangci/golangci-lint-action@v3 + uses: golangci/golangci-lint-action@v9 with: - version: v1.50 + version: v2.9 diff --git a/.golangci.yml b/.golangci.yml index 7a1516a..2aa0fa1 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -1,25 +1,31 @@ -issues: - exclude-rules: - - linters: - - testpackage - - gochecknoglobals - path: '(.+)_test\.go' +version: '2' linters: - enable-all: true + default: all disable: - # deprecated - - maligned - - scopelint - - interfacer - - golint - - exhaustivestruct - - nosnakecase - - structcheck - - deadcode - - varcheck - - ifshort - # unused - exhaustruct - - nlreturn -run: - timeout: 3m \ No newline at end of file + - depguard + - wsl + - testpackage + settings: + gocritic: + enable-all: true + settings: + unnamedResult: + checkExported: true + errcheck: + check-type-assertions: true + check-blank: false + disable-default-exclusions: false + exclude-functions: + - (*os.File).Close + - (io.Closer).Close + +issues: + max-issues-per-linter: 0 + max-same-issues: 0 +formatters: + enable: + - gci + - gofmt + - gofumpt + - goimports diff --git a/LICENSE b/LICENSE index e8c9e4c..80384d3 100644 --- a/LICENSE +++ b/LICENSE @@ -1,6 +1,6 @@ MIT License -Copyright (c) 2023 David Newhall II +Copyright (c) 2019-2026 David Newhall II Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal diff --git a/concurrency_test.go b/concurrency_test.go new file mode 100644 index 0000000..ab126d7 --- /dev/null +++ b/concurrency_test.go @@ -0,0 +1,128 @@ +package subscribe + +import ( + "errors" + "path/filepath" + "strconv" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestEventsConcurrentAccess(t *testing.T) { + t.Parallel() + + events := &Events{Map: make(map[string]*Rules)} + require.NoError(t, events.New("event", nil)) + + var waitGroup sync.WaitGroup + + for i := range 32 { + waitGroup.Go(func() { + rule := "rule_" + strconv.Itoa(i) + for j := range 200 { + events.RuleSetI("event", rule, j) + events.RuleSetS("event", rule, rule) + events.RuleSetD("event", rule, time.Duration(j)*time.Millisecond) + events.RuleSetT("event", rule, time.Now().Add(time.Duration(j)*time.Second)) + + err := events.Pause("event", time.Millisecond) + if err != nil { + t.Errorf("pause failed: %v", err) + + return + } + + err = events.UnPause("event") + if err != nil { + t.Errorf("unpause failed: %v", err) + + return + } + + events.IsPaused("event") + events.PauseTime("event") + events.RuleGetI("event", rule) + events.RuleGetS("event", rule) + events.RuleGetD("event", rule) + events.RuleGetT("event", rule) + events.RuleDelAll("event", rule) + } + }) + } + + waitGroup.Wait() +} + +func TestSubscribeConcurrentAccess(t *testing.T) { + t.Parallel() + + stateFile := filepath.Join(t.TempDir(), "state.json") + sub, err := GetDB(stateFile) + require.NoError(t, err) + require.NoError(t, sub.Events.New("evt", nil)) + + var waitGroup sync.WaitGroup + + for i := range 20 { + waitGroup.Go(func() { + contact := "contact_" + strconv.Itoa(i%5) + runSubscribeOps(t, sub, contact) + }) + } + + waitGroup.Wait() +} + +func runSubscribeOps(t *testing.T, sub *Subscribe, contact string) { + t.Helper() + + for j := range 100 { + admin := j%2 == 0 + ignored := j%3 == 0 + subscriber := sub.CreateSub(contact, "api", admin, ignored) + + if subscriber == nil { + t.Error("subscriber should not be nil") + + return + } + + err := subscriber.Subscribe("evt") + + if err != nil && !errors.Is(err, ErrEventExists) { + t.Errorf("subscribe failed: %v", err) + + return + } + + err = subscriber.Events.Pause("evt", time.Millisecond) + if err != nil { + t.Errorf("pause failed: %v", err) + + return + } + + err = subscriber.Events.UnPause("evt") + if err != nil { + t.Errorf("unpause failed: %v", err) + + return + } + + _, _ = sub.GetSubscriber(contact, "api") + sub.GetAdmins() + sub.GetIgnored() + sub.GetSubscribers("evt") + _, _ = sub.StateGetJSON() + + err = sub.StateFileSave() + if err != nil { + t.Errorf("state save failed: %v", err) + + return + } + } +} diff --git a/database.go b/database.go index b67178f..9460988 100644 --- a/database.go +++ b/database.go @@ -1,9 +1,12 @@ +// Package subscribe provides a subscription management system. package subscribe import ( "encoding/json" "fmt" + "maps" "os" + "time" ) /************************ @@ -19,48 +22,82 @@ func GetDB(stateFile string) (*Subscribe, error) { Subscribers: make([]*Subscriber, 0), } - return sub, sub.StateFileLoad() + err := sub.StateFileLoad() + if err != nil { + return nil, err + } + + return sub, nil } // StateFileLoad data from a json file. func (s *Subscribe) StateFileLoad() error { - if s.stateFile == "" { + s.mu.RLock() + stateFile := s.stateFile + s.mu.RUnlock() + + if stateFile == "" { return nil } - if buf, err := os.ReadFile(s.stateFile); os.IsNotExist(err) { + // #nosec G304 -- state file path is user-configured on purpose. + buf, err := os.ReadFile(stateFile) + if os.IsNotExist(err) { return s.StateFileSave() - } else if err != nil { - return fmt.Errorf("file problem: %w", err) - } else if err := json.Unmarshal(buf, s); err != nil { - return fmt.Errorf("json problem: %w", err) } + if err != nil { + return fmt.Errorf("failed reading state file: %w", err) + } + + loaded := new(Subscribe) + + err = json.Unmarshal(buf, loaded) + if err != nil { + return fmt.Errorf("failed decoding state file: %w", err) + } + + normalizeLoadedState(loaded) + + s.mu.Lock() + s.EnableAPIs = loaded.EnableAPIs + s.Events = loaded.Events + s.Subscribers = loaded.Subscribers + s.mu.Unlock() + return nil } // StateGetJSON returns the state data in json format. func (s *Subscribe) StateGetJSON() (string, error) { - s.Events.RLock() - defer s.Events.RUnlock() + snapshot := s.snapshot() - b, err := json.Marshal(s) + b, err := json.Marshal(snapshot) return string(b), err } // StateFileSave writes out the state file. func (s *Subscribe) StateFileSave() error { - if s.stateFile == "" { + const stateFileMode = 0o600 + + s.mu.RLock() + stateFile := s.stateFile + s.mu.RUnlock() + + if stateFile == "" { return nil } - s.Events.RLock() - defer s.Events.RUnlock() + snapshot := s.snapshot() - if buf, err := json.Marshal(s); err != nil { + buf, err := json.Marshal(snapshot) + if err != nil { return fmt.Errorf("marshaling json: %w", err) - } else if err = os.WriteFile(s.stateFile, buf, 0o600); err != nil { //nolint:gomnd + } + + err = os.WriteFile(stateFile, buf, stateFileMode) + if err != nil { return fmt.Errorf("writing file: %w", err) } @@ -69,12 +106,137 @@ func (s *Subscribe) StateFileSave() error { // StateFileRelocate writes the state file to a new location. func (s *Subscribe) StateFileRelocate(newPath string) error { - s.stateFile, newPath = newPath, s.stateFile // swap places + s.mu.Lock() + oldPath := s.stateFile + s.stateFile = newPath + s.mu.Unlock() + + err := s.StateFileLoad() + if err != nil { + s.mu.Lock() + s.stateFile = oldPath + s.mu.Unlock() + } + + return err +} - if err := s.StateFileLoad(); err != nil { - s.stateFile = newPath // got an error, put it back. - return err +func normalizeLoadedState(loaded *Subscribe) { + if loaded.EnableAPIs == nil { + loaded.EnableAPIs = make([]string, 0) } - return nil + if loaded.Events == nil { + loaded.Events = &Events{Map: make(map[string]*Rules)} + } else { + normalizeEvents(loaded.Events) + } + + if loaded.Subscribers == nil { + loaded.Subscribers = make([]*Subscriber, 0) + } + + for _, sub := range loaded.Subscribers { + if sub == nil { + continue + } + + if sub.Events == nil { + sub.Events = &Events{Map: make(map[string]*Rules)} + + continue + } + + normalizeEvents(sub.Events) + } +} + +func normalizeEvents(events *Events) { + if events.Map == nil { + events.Map = make(map[string]*Rules) + } + + for key, rules := range events.Map { + if rules == nil { + events.Map[key] = &Rules{ + D: make(map[string]time.Duration), + I: make(map[string]int), + S: make(map[string]string), + T: make(map[string]time.Time), + } + + continue + } + + if rules.D == nil { + rules.D = make(map[string]time.Duration) + } + + if rules.I == nil { + rules.I = make(map[string]int) + } + + if rules.S == nil { + rules.S = make(map[string]string) + } + + if rules.T == nil { + rules.T = make(map[string]time.Time) + } + } +} + +func (s *Subscribe) snapshot() *Subscribe { + s.mu.RLock() + defer s.mu.RUnlock() + + out := &Subscribe{ + EnableAPIs: append(make([]string, 0, len(s.EnableAPIs)), s.EnableAPIs...), + Events: snapshotEvents(s.Events), + Subscribers: make([]*Subscriber, 0, len(s.Subscribers)), + } + + for _, sub := range s.Subscribers { + out.Subscribers = append(out.Subscribers, snapshotSubscriber(sub)) + } + + return out +} + +func snapshotSubscriber(sub *Subscriber) *Subscriber { + if sub == nil { + return nil + } + + out := &Subscriber{ + ID: sub.ID, + API: sub.API, + Contact: sub.Contact, + Events: snapshotEvents(sub.Events), + Admin: sub.Admin, + Ignored: sub.Ignored, + } + + if sub.Meta != nil { + out.Meta = make(map[string]any, len(sub.Meta)) + maps.Copy(out.Meta, sub.Meta) + } + + return out +} + +func snapshotEvents(events *Events) *Events { + if events == nil { + return &Events{Map: make(map[string]*Rules)} + } + + events.mu.RLock() + defer events.mu.RUnlock() + + out := &Events{Map: make(map[string]*Rules, len(events.Map))} + for event, rules := range events.Map { + out.Map[event] = cloneRules(rules) + } + + return out } diff --git a/database_test.go b/database_test.go index fa6d155..a81bc28 100644 --- a/database_test.go +++ b/database_test.go @@ -6,100 +6,127 @@ import ( "testing" "github.com/stretchr/testify/assert" -) - -var ( - testFile = filepath.Join(os.TempDir(), "this_is_a_testfile_for_subtscribe_test.go.json") - testFile2 = filepath.Join(os.TempDir(), "this_is_a_testfile_for_subtscribe_test2.go.json") - testFile4 = filepath.Join(os.TempDir(), "this_is_a_testfile_for_subtscribe_test4.go.json") + "github.com/stretchr/testify/require" ) func TestGetDB(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub, err := GetDB("") - assert.Nil(err, "getting an empty db must produce no error") + require.NoError(t, err, "getting an empty db must produce no error") json, err := sub.StateGetJSON() - assert.EqualValues(`{"enabledApis":[],"events":{"eventsMap":{}},"subscribers":[]}`, json, + assertions.JSONEq(`{"enabledApis":[],"events":{"eventsMap":{}},"subscribers":[]}`, json, "the initial state must be empty") - assert.Nil(err, "getting an empty state must produce no error") + require.NoError(t, err, "getting an empty state must produce no error") } func TestStateFileLoad(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) + testFile := filepath.Join(t.TempDir(), "state.json") // test with good data. testJSON := `{"enabledApis":[],"events":{"eventsMap":{}},"subscribers":[{"id":0,"meta":null,"api":` + `"http","contact":"testUser","events":{"eventsMap":{}},"isAdmin":false,"ignored":false}]}` - assert.Nil(os.WriteFile(testFile, []byte(testJSON), 0o600), "problem writing test file") + require.NoError(t, os.WriteFile(testFile, []byte(testJSON), 0o600), "problem writing test file") sub, err := GetDB(testFile) - assert.Nil(err, "there must be no error loading the state file") + require.NoError(t, err, "there must be no error loading the state file") json, err := sub.StateGetJSON() - assert.Nil(err, "there must be no error getting the state data") - assert.EqualValues(testJSON, json) + require.NoError(t, err, "there must be no error getting the state data") + assertions.JSONEq(testJSON, json) // Test missing file. - assert.Nil(os.RemoveAll(testFile), "problem removing test file") + require.NoError(t, os.RemoveAll(testFile), "problem removing test file") _, err = GetDB(testFile) - assert.Nil(err, "there must be no error when the state file is missing") + require.NoError(t, err, "there must be no error when the state file is missing") + // #nosec G304 -- test controls this temporary file path. data, err := os.ReadFile(testFile) - assert.Nil(err, "error reading test file") + require.NoError(t, err, "error reading test file") - assert.EqualValues(`{"enabledApis":[],"events":{"eventsMap":{}},"subscribers":[]}`, data, + assertions.JSONEq(`{"enabledApis":[],"events":{"eventsMap":{}},"subscribers":[]}`, string(data), "the initial state file must be empty") // Test uncreatable file. _, err = GetDB("/tmp/xxx/yyy/zzz/aaa/bbb/this_file_dont_exist") - assert.NotNil(err, "there must be an error when the state cannot be created") + require.Error(t, err, "there must be an error when the state cannot be created") - // Test unreadable file. - _, err = GetDB("/etc/sudoers") - assert.NotNil(err, "there must be an error when the state cannot be read") + // Test unreadable path (directory). + _, err = GetDB(t.TempDir()) + require.Error(t, err, "there must be an error when the state path is not readable as a file") // Test bad data. err = os.WriteFile(testFile, []byte("this aint good json}}"), 0o600) - assert.Nil(err, "problem writing test file") + require.NoError(t, err, "problem writing test file") _, err = GetDB(testFile) - assert.NotNil(err, "there must be an error when the state file is corrupt") + require.Error(t, err, "there must be an error when the state file is corrupt") } func TestStateFileSave(t *testing.T) { t.Parallel() - assert := assert.New(t) - assert.Nil(os.RemoveAll(testFile2), "problem removing test file") - sub, err := GetDB(testFile2) - assert.Nil(err, "there must be no error creating the initial state file") - assert.Nil(sub.StateFileSave(), "there must be no error saving the state file") + + testFile := filepath.Join(t.TempDir(), "state-save.json") + sub, err := GetDB(testFile) + require.NoError(t, err, "there must be no error creating the initial state file") + require.NoError(t, sub.StateFileSave(), "there must be no error saving the state file") sub, err = GetDB("") - assert.Nil(err, "there must be no error when the state file does not exist") - assert.Nil(sub.StateFileSave(), "there must be no error when the state file does not exist") + require.NoError(t, err, "there must be no error when the state file does not exist") + require.NoError(t, sub.StateFileSave(), "there must be no error when the state file does not exist") } func TestStateFileRelocate(t *testing.T) { t.Parallel() - assert := assert.New(t) - assert.Nil(os.RemoveAll(testFile4), "problem removing test file") + assertions := assert.New(t) + testFile4 := filepath.Join(t.TempDir(), "state-relocate.json") sub, err := GetDB("") - assert.Nil(err, "there must be no error when the state file does not exist") - assert.Nil(sub.StateFileSave(), "there must be no error when the state file does not exist") + require.NoError(t, err, "there must be no error when the state file does not exist") + require.NoError(t, sub.StateFileSave(), "there must be no error when the state file does not exist") err = sub.StateFileRelocate(testFile4) - assert.EqualValues(testFile4, sub.stateFile, "the path was not to the new value") - assert.Nil(err, "there must be no error creating the initial state file") + assertions.Equal(testFile4, sub.stateFile, "the path was not to the new value") + require.NoError(t, err, "there must be no error creating the initial state file") + + err = sub.StateFileRelocate(t.TempDir()) + require.Error(t, err, "there must be an error trying to write a file as a folder") + assertions.Equal(testFile4, sub.stateFile, "the path was not changed back to the previous value") +} + +func TestStateGetJSONMarshalError(t *testing.T) { + t.Parallel() + + sub := &Subscribe{ + Events: &Events{Map: make(map[string]*Rules)}, + Subscribers: make([]*Subscriber, 0), + } + + subscriber := sub.CreateSub("marshal", "api", false, false) + subscriber.Meta = map[string]any{"bad": make(chan int)} + + _, err := sub.StateGetJSON() + require.Error(t, err, "unsupported meta values should fail marshaling") +} + +func TestStateFileSaveMarshalError(t *testing.T) { + t.Parallel() + + stateFile := filepath.Join(t.TempDir(), "state-save-marshal.json") + sub, err := GetDB(stateFile) + require.NoError(t, err) + + subscriber := sub.CreateSub("marshal", "api", false, false) + subscriber.Meta = map[string]any{"bad": make(chan int)} - err = sub.StateFileRelocate(os.TempDir()) - assert.NotNil(err, "there must be an error trying to write a file as a tmp folder") - assert.EqualValues(testFile4, sub.stateFile, "the path was not changed back to the previous value") + err = sub.StateFileSave() + require.Error(t, err) + assert.ErrorContains(t, err, "marshaling json") } diff --git a/events.go b/events.go index 40ceaea..a3c06ce 100644 --- a/events.go +++ b/events.go @@ -1,6 +1,7 @@ package subscribe import ( + "maps" "sort" "strings" "time" @@ -12,10 +13,10 @@ import ( // Names returns all the configured event names. func (e *Events) Names() []string { - e.RLock() - defer e.RUnlock() + e.mu.RLock() + defer e.mu.RUnlock() - names := []string{} + names := make([]string, 0, len(e.Map)) for name := range e.Map { names = append(names, name) @@ -28,16 +29,16 @@ func (e *Events) Names() []string { // Len returns the number of configured events. func (e *Events) Len() int { - e.RLock() - defer e.RUnlock() + e.mu.RLock() + defer e.mu.RUnlock() return len(e.Map) } -// Name finds an event case insentively. +// Name finds an event case insensitively. func (e *Events) Name(event string) string { - e.RLock() - defer e.RUnlock() + e.mu.RLock() + defer e.mu.RUnlock() if _, ok := e.Map[event]; ok { return event @@ -54,8 +55,8 @@ func (e *Events) Name(event string) string { // Exists returns true if an event exists. func (e *Events) Exists(event string) bool { - e.RLock() - defer e.RUnlock() + e.mu.RLock() + defer e.mu.RUnlock() if _, ok := e.Map[event]; ok { return true @@ -66,34 +67,14 @@ func (e *Events) Exists(event string) bool { // New adds an event. func (e *Events) New(event string, rules *Rules) error { - e.Lock() - defer e.Unlock() + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; ok { return ErrEventExists } - if rules == nil { - rules = &Rules{} - } - - if rules.D == nil { - rules.D = make(map[string]time.Duration) - } - - if rules.I == nil { - rules.I = make(map[string]int) - } - - if rules.S == nil { - rules.S = make(map[string]string) - } - - if rules.T == nil { - rules.T = make(map[string]time.Time) - } - - e.Map[event] = rules + e.Map[event] = cloneRules(rules) return nil } @@ -107,8 +88,8 @@ func (e *Events) UnPause(event string) error { // Pause (or unpause with 0 duration) a subscriber's event subscription. // Returns an error only if the event subscription is not found. func (e *Events) Pause(event string, duration time.Duration) error { - e.RLock() - defer e.RUnlock() + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; !ok { return ErrEventNotFound @@ -119,11 +100,11 @@ func (e *Events) Pause(event string, duration time.Duration) error { return nil } -// IsPaused returns true if the event's notifications are pasued. +// IsPaused returns true if the event's notifications are paused. // Returns true if the event subscription does not exist. func (e *Events) IsPaused(event string) bool { - e.RLock() - defer e.RUnlock() + e.mu.RLock() + defer e.mu.RUnlock() info, ok := e.Map[event] if !ok { @@ -135,8 +116,8 @@ func (e *Events) IsPaused(event string) bool { // PauseTime returns the pause time for an event. func (e *Events) PauseTime(event string) time.Time { - e.RLock() - defer e.RUnlock() + e.mu.RLock() + defer e.mu.RUnlock() info, ok := e.Map[event] if !ok { @@ -148,92 +129,76 @@ func (e *Events) PauseTime(event string) time.Time { // Remove deletes an event. func (e *Events) Remove(event string) { - e.Lock() - defer e.Unlock() + e.mu.Lock() + defer e.mu.Unlock() delete(e.Map, event) } // RuleGetD returns a Duration rule. func (e *Events) RuleGetD(event, rule string) (time.Duration, bool) { - e.RLock() - defer e.RUnlock() + e.mu.RLock() + defer e.mu.RUnlock() - r, ok := e.Map[event] - if !ok || r == nil { + rules, found := e.Map[event] + if !found || rules == nil { return 0, false } - for n, v := range r.D { - if n == rule { - return v, true - } - } + val, found := rules.D[rule] - return 0, false + return val, found } // RuleGetI returns an integer rule. func (e *Events) RuleGetI(event, rule string) (int, bool) { - e.RLock() - defer e.RUnlock() + e.mu.RLock() + defer e.mu.RUnlock() - r, ok := e.Map[event] - if !ok || r == nil { + rules, found := e.Map[event] + if !found || rules == nil { return 0, false } - for n, v := range r.I { - if n == rule { - return v, true - } - } + val, found := rules.I[rule] - return 0, false + return val, found } // RuleGetS returns a string rule. func (e *Events) RuleGetS(event, rule string) (string, bool) { - e.RLock() - defer e.RUnlock() + e.mu.RLock() + defer e.mu.RUnlock() - r, ok := e.Map[event] - if !ok || r == nil { + rules, found := e.Map[event] + if !found || rules == nil { return "", false } - for n, v := range r.S { - if n == rule { - return v, true - } - } + val, found := rules.S[rule] - return "", false + return val, found } // RuleGetT returns a Time rule. func (e *Events) RuleGetT(event, rule string) (time.Time, bool) { - e.RLock() - defer e.RUnlock() + e.mu.RLock() + defer e.mu.RUnlock() - r, ok := e.Map[event] - if !ok || r == nil { + rules, found := e.Map[event] + if !found || rules == nil { return time.Now(), false } - for n, v := range r.T { - if n == rule { - return v, true - } - } + val, found := rules.T[rule] - return time.Now(), false + return val, found } // RuleSetD updates or sets a Duration rule. func (e *Events) RuleSetD(event, rule string, val time.Duration) { - e.Lock() - defer e.Unlock() + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; !ok { return @@ -248,8 +213,8 @@ func (e *Events) RuleSetD(event, rule string, val time.Duration) { // RuleSetI updates or sets an integer rule. func (e *Events) RuleSetI(event, rule string, val int) { - e.Lock() - defer e.Unlock() + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; !ok { return @@ -263,9 +228,9 @@ func (e *Events) RuleSetI(event, rule string, val int) { } // RuleSetS updates or sets a string rule. -func (e *Events) RuleSetS(event, rule string, val string) { - e.Lock() - defer e.Unlock() +func (e *Events) RuleSetS(event, rule, val string) { + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; !ok { return @@ -280,8 +245,8 @@ func (e *Events) RuleSetS(event, rule string, val string) { // RuleSetT updates or sets a Time rule. func (e *Events) RuleSetT(event, rule string, val time.Time) { - e.Lock() - defer e.Unlock() + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; !ok { return @@ -296,8 +261,8 @@ func (e *Events) RuleSetT(event, rule string, val time.Time) { // RuleDelD deletes a Duration rule. func (e *Events) RuleDelD(event, rule string) { - e.Lock() - defer e.Unlock() + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; !ok || e.Map[event].D == nil { return @@ -308,8 +273,8 @@ func (e *Events) RuleDelD(event, rule string) { // RuleDelI deletes an integer rule. func (e *Events) RuleDelI(event, rule string) { - e.Lock() - defer e.Unlock() + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; !ok || e.Map[event].I == nil { return @@ -320,8 +285,8 @@ func (e *Events) RuleDelI(event, rule string) { // RuleDelS deletes a string rule. func (e *Events) RuleDelS(event, rule string) { - e.Lock() - defer e.Unlock() + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; !ok || e.Map[event].S == nil { return @@ -332,8 +297,8 @@ func (e *Events) RuleDelS(event, rule string) { // RuleDelT deletes a Time rule. func (e *Events) RuleDelT(event, rule string) { - e.Lock() - defer e.Unlock() + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; !ok || e.Map[event].T == nil { return @@ -344,8 +309,8 @@ func (e *Events) RuleDelT(event, rule string) { // RuleDelAll deletes rules of any type with a specific name. func (e *Events) RuleDelAll(event, rule string) { - e.Lock() - defer e.Unlock() + e.mu.Lock() + defer e.mu.Unlock() if _, ok := e.Map[event]; !ok { return @@ -356,3 +321,29 @@ func (e *Events) RuleDelAll(event, rule string) { delete(e.Map[event].S, rule) delete(e.Map[event].T, rule) } + +func cloneRules(rules *Rules) *Rules { + if rules == nil { + return &Rules{ + D: make(map[string]time.Duration), + I: make(map[string]int), + S: make(map[string]string), + T: make(map[string]time.Time), + } + } + + cloned := &Rules{ + Pause: rules.Pause, + D: make(map[string]time.Duration, len(rules.D)), + I: make(map[string]int, len(rules.I)), + S: make(map[string]string, len(rules.S)), + T: make(map[string]time.Time, len(rules.T)), + } + + maps.Copy(cloned.D, rules.D) + maps.Copy(cloned.I, rules.I) + maps.Copy(cloned.S, rules.S) + maps.Copy(cloned.T, rules.T) + + return cloned +} diff --git a/events_test.go b/events_test.go index 992087e..3493056 100644 --- a/events_test.go +++ b/events_test.go @@ -1,66 +1,231 @@ package subscribe import ( - "os" "path/filepath" + "sort" "testing" + "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -var testFile3 = filepath.Join(os.TempDir(), "this_is_a_testfile_for_subtscribe_test3.go.json") - func TestNew(t *testing.T) { t.Parallel() - assert := assert.New(t) - sub, err := GetDB(testFile3) + assertions := assert.New(t) + stateFile := filepath.Join(t.TempDir(), "events_state.json") + sub, err := GetDB(stateFile) - assert.NotNil(sub.Events.Names(), "the events slice must not be nil") - assert.Nil(err, "getting a db must produce no error") - assert.EqualValues(0, len(sub.Events.Names()), "event count must be 0 since none have been added") - assert.Nil(sub.Events.New("event_test", nil)) - assert.EqualValues(1, len(sub.Events.Names()), "event count must be 1 since 1 was added") + assertions.NotNil(sub.Events.Names(), "the events slice must not be nil") + require.NoError(t, err, "getting db must produce no error") + assertions.Empty(sub.Events.Names(), "event count must be 0 since none have been added") + require.NoError(t, sub.Events.New("event_test", nil)) + assertions.Len(sub.Events.Names(), 1, "event count must be 1 since 1 was added") } func TestGetEvent(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub, err := GetDB("") - assert.NotNil(sub.Events.Names(), "the events map must not be nil") - assert.Nil(err, "getting a db must produce no error") - assert.Nil(sub.Events.New("event_test", nil)) - assert.True(sub.Events.Exists("event_test"), "this event exists so the method must return true") - assert.EqualValues(1, len(sub.Events.Names()), "event count must be 1 since 1 was added") - assert.False(sub.Events.Exists("missing_event"), "this event does not exists") + assertions.NotNil(sub.Events.Names(), "the events map must not be nil") + require.NoError(t, err, "getting db must produce no error") + require.NoError(t, sub.Events.New("event_test", nil)) + assertions.True(sub.Events.Exists("event_test"), "this event exists so the method must return true") + assertions.Len(sub.Events.Names(), 1, "event count must be 1 since 1 was added") + assertions.False(sub.Events.Exists("missing_event"), "this event does not exist") } func TestNewEvent(t *testing.T) { t.Parallel() - a := assert.New(t) + assertions := assert.New(t) sub := &Subscribe{Events: &Events{Map: make(map[string]*Rules)}} - a.Nil(sub.Events.New("event_test", nil)) - a.NotNil(sub.Events.Map["event_test"], "the event rules map must not be nil") + require.NoError(t, sub.Events.New("event_test", nil)) + assertions.NotNil(sub.Events.Map["event_test"], "the event rules map must not be nil") + assertions.NotNil(sub.Events.Map["event_test"].D, "duration map must be initialized") + assertions.NotNil(sub.Events.Map["event_test"].I, "integer map must be initialized") + assertions.NotNil(sub.Events.Map["event_test"].S, "string map must be initialized") + assertions.NotNil(sub.Events.Map["event_test"].T, "time map must be initialized") } func TestRemoveEvent(t *testing.T) { t.Parallel() - assert := assert.New(t) sub := &Subscribe{Events: &Events{Map: make(map[string]*Rules)}} - sub.Events.Remove("no_event") // Make two events to remove. sub.Events.Map["some_event"] = nil sub.Events.Map["some_event2"] = nil - // Subscribe a user to one of them. - s := sub.CreateSub("test_contact", "api", true, false) // XXX: count them? - assert.Nil(s.Subscribe("some_event2")) + // Subscribe asert user to one of them. + subscriber := sub.CreateSub("test_contact", "api", true, false) + require.NoError(t, subscriber.Subscribe("some_event2")) sub.EventRemove("some_event2") sub.EventRemove("some_event") + + assert.False(t, sub.Events.Exists("some_event"), "global event should be removed") + assert.False(t, sub.Events.Exists("some_event2"), "global event should be removed") + assert.False(t, subscriber.Events.Exists("some_event2"), "subscription event should be removed") +} + +func TestEventsName(t *testing.T) { + t.Parallel() + + events := &Events{Map: make(map[string]*Rules)} + require.NoError(t, events.New("CaseSensitive", nil)) + require.NoError(t, events.New("other", nil)) + + assert.Equal(t, "CaseSensitive", events.Name("CaseSensitive")) + assert.Equal(t, "CaseSensitive", events.Name("casesensitive")) + assert.Empty(t, events.Name("missing")) +} + +func TestEventsPauseTimeAndLen(t *testing.T) { + t.Parallel() + + events := &Events{Map: make(map[string]*Rules)} + assert.Equal(t, 0, events.Len()) + assert.Equal(t, time.Time{}, events.PauseTime("missing")) + + require.NoError(t, events.New("pause_test", nil)) + require.NoError(t, events.Pause("pause_test", 2*time.Minute)) + assert.Equal(t, 1, events.Len()) + assert.WithinDuration(t, time.Now().Add(2*time.Minute), events.PauseTime("pause_test"), 2*time.Second) +} + +func TestEventsRuleLifecycleAllTypes(t *testing.T) { + t.Parallel() + + events := &Events{Map: make(map[string]*Rules)} + require.NoError(t, events.New("event", nil)) + + when := time.Now().UTC().Round(time.Second) + + events.RuleSetD("event", "d", 3*time.Minute) + events.RuleSetI("event", "i", 55) + events.RuleSetS("event", "s", "value") + events.RuleSetT("event", "t", when) + + d, found := events.RuleGetD("event", "d") + require.True(t, found) + assert.Equal(t, 3*time.Minute, d) + + i, found := events.RuleGetI("event", "i") + require.True(t, found) + assert.Equal(t, 55, i) + + s, found := events.RuleGetS("event", "s") + require.True(t, found) + assert.Equal(t, "value", s) + + gotTime, found := events.RuleGetT("event", "t") + require.True(t, found) + assert.Equal(t, when, gotTime) + + events.RuleDelD("event", "d") + events.RuleDelI("event", "i") + events.RuleDelS("event", "s") + events.RuleDelT("event", "t") + + _, found = events.RuleGetD("event", "d") + assert.False(t, found) + _, found = events.RuleGetI("event", "i") + assert.False(t, found) + _, found = events.RuleGetS("event", "s") + assert.False(t, found) + _, found = events.RuleGetT("event", "t") + assert.False(t, found) +} + +func TestEventsRuleDelAll(t *testing.T) { + t.Parallel() + + events := &Events{Map: make(map[string]*Rules)} + require.NoError(t, events.New("event", nil)) + + events.RuleSetD("event", "shared", time.Second) + events.RuleSetI("event", "shared", 1) + events.RuleSetS("event", "shared", "s") + events.RuleSetT("event", "shared", time.Now()) + events.RuleDelAll("event", "shared") + + _, found := events.RuleGetD("event", "shared") + assert.False(t, found) + _, found = events.RuleGetI("event", "shared") + assert.False(t, found) + _, found = events.RuleGetS("event", "shared") + assert.False(t, found) + _, found = events.RuleGetT("event", "shared") + assert.False(t, found) +} + +func TestEventsRuleGetMissingEvent(t *testing.T) { + t.Parallel() + + events := &Events{Map: make(map[string]*Rules)} + events.Map["bad"] = nil + + _, found := events.RuleGetD("missing", "r") + assert.False(t, found) + _, found = events.RuleGetI("missing", "r") + assert.False(t, found) + _, found = events.RuleGetS("missing", "r") + assert.False(t, found) + _, found = events.RuleGetT("missing", "r") + assert.False(t, found) + + _, found = events.RuleGetD("bad", "r") + assert.False(t, found) + _, found = events.RuleGetI("bad", "r") + assert.False(t, found) + _, found = events.RuleGetS("bad", "r") + assert.False(t, found) + _, found = events.RuleGetT("bad", "r") + assert.False(t, found) +} + +func TestEventsNewClonesRules(t *testing.T) { + t.Parallel() + + events := &Events{Map: make(map[string]*Rules)} + external := &Rules{ + D: map[string]time.Duration{"a": time.Second}, + I: map[string]int{"b": 2}, + S: map[string]string{"c": "value"}, + T: map[string]time.Time{"d": time.Now()}, + } + require.NoError(t, events.New("event", external)) + + external.D["a"] = 4 * time.Second + external.I["b"] = 42 + external.S["c"] = "changed" + external.T["d"] = time.Now().Add(time.Hour) + + d, _ := events.RuleGetD("event", "a") + i, _ := events.RuleGetI("event", "b") + s, _ := events.RuleGetS("event", "c") + ts, _ := events.RuleGetT("event", "d") + + assert.Equal(t, time.Second, d) + assert.Equal(t, 2, i) + assert.Equal(t, "value", s) + assert.False(t, ts.IsZero()) +} + +func TestEventsNamesSorted(t *testing.T) { + t.Parallel() + + events := &Events{Map: make(map[string]*Rules)} + require.NoError(t, events.New("c", nil)) + require.NoError(t, events.New("a", nil)) + require.NoError(t, events.New("b", nil)) + + names := events.Names() + expected := []string{"a", "b", "c"} + sort.Strings(expected) + assert.Equal(t, expected, names) } diff --git a/go.mod b/go.mod index 1de938b..b8e5d3f 100644 --- a/go.mod +++ b/go.mod @@ -1,8 +1,10 @@ module golift.io/subscribe -go 1.17 +go 1.25.6 -require github.com/stretchr/testify v1.8.4 +toolchain go1.26.0 + +require github.com/stretchr/testify v1.11.1 require ( github.com/davecgh/go-spew v1.1.1 // indirect diff --git a/go.sum b/go.sum index da9610b..c4c1710 100644 --- a/go.sum +++ b/go.sum @@ -1,27 +1,10 @@ -github.com/davecgh/go-spew v1.1.0 h1:ZDRjVQ15GmhC3fiQ8ni8+OwkZQO4DARzQgrnXU1Liz8= -github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= -github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= -github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= -github.com/stretchr/testify v1.6.1 h1:hDPOHmpOpP40lSULcqw7IrRb/u7w6RpDC9399XyoNd0= -github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= -github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= -github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= -github.com/stretchr/testify v1.8.1 h1:w7B6lhMri9wdJUVmEZPGGhZzrYTPvgJArz7wNPgYKsk= -github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= -github.com/stretchr/testify v1.8.2 h1:+h33VjcLVPDHtOdpUCuF+7gSuG3yGIftsP1YvFihtJ8= -github.com/stretchr/testify v1.8.2/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= -github.com/stretchr/testify v1.8.3 h1:RP3t2pwF7cMEbC1dqtB6poj3niw/9gnV4Cjg5oW5gtY= -github.com/stretchr/testify v1.8.3/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= -github.com/stretchr/testify v1.8.4 h1:CcVxjf3Q8PM0mHUKJCdn+eZZtm5yQwehR5yeSVQQcUk= -github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c h1:dUUwHk2QECo/6vqA44rthZ8ie2QXMNeKRTHCNY2nXvo= -gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/subscriber.go b/subscriber.go index da04aba..f8bb4ef 100644 --- a/subscriber.go +++ b/subscriber.go @@ -6,6 +6,9 @@ package subscribe // CreateSub creates or updates a subscriber. func (s *Subscribe) CreateSub(contact, api string, admin, ignore bool) *Subscriber { + s.mu.Lock() + defer s.mu.Unlock() + for i := range s.Subscribers { if contact == s.Subscribers[i].Contact && api == s.Subscribers[i].API { s.Subscribers[i].Admin = admin @@ -28,17 +31,21 @@ func (s *Subscribe) CreateSub(contact, api string, admin, ignore bool) *Subscrib return s.Subscribers[len(s.Subscribers)-1] } +// CreateSubWithID creates or updates a subscriber with a given ID. func (s *Subscribe) CreateSubWithID(subID int64, contact, api string, admin, ignore bool) *Subscriber { if subID == 0 { return nil } - for idx := range s.Subscribers { - if subID == s.Subscribers[idx].ID && api == s.Subscribers[idx].API { - s.Subscribers[idx].Admin = admin - s.Subscribers[idx].Ignored = ignore + s.mu.Lock() + defer s.mu.Unlock() + + for i := range s.Subscribers { + if subID == s.Subscribers[i].ID && api == s.Subscribers[i].API { + s.Subscribers[i].Admin = admin + s.Subscribers[i].Ignored = ignore // Already exists, return it. - return s.Subscribers[idx] + return s.Subscribers[i] } } @@ -61,6 +68,9 @@ func (s *Subscribe) CreateSubWithID(subID int64, contact, api string, admin, ign // GetSubscriber gets a subscriber based on their contact info. func (s *Subscribe) GetSubscriber(contact, api string) (*Subscriber, error) { + s.mu.RLock() + defer s.mu.RUnlock() + for _, sub := range s.Subscribers { if sub.Contact == contact && sub.API == api { return sub, nil @@ -76,6 +86,9 @@ func (s *Subscribe) GetSubscriberByID(subID int64, api string) (*Subscriber, err return nil, ErrSubscriberNotFound } + s.mu.RLock() + defer s.mu.RUnlock() + for _, sub := range s.Subscribers { if sub.ID == subID && sub.API == api { return sub, nil @@ -87,11 +100,14 @@ func (s *Subscribe) GetSubscriberByID(subID int64, api string) (*Subscriber, err // GetAdmins returns a list of subscribed admins. func (s *Subscribe) GetAdmins() []*Subscriber { - var subs []*Subscriber + s.mu.RLock() + defer s.mu.RUnlock() - for _, sub := range s.Subscribers { - if sub.Admin { - subs = append(subs, sub) + subs := make([]*Subscriber, 0, len(s.Subscribers)) + + for idx := range s.Subscribers { + if s.Subscribers[idx].Admin { + subs = append(subs, s.Subscribers[idx]) } } @@ -100,11 +116,14 @@ func (s *Subscribe) GetAdmins() []*Subscriber { // GetIgnored returns a list of ignored subscribers. func (s *Subscribe) GetIgnored() []*Subscriber { - var subs []*Subscriber + s.mu.RLock() + defer s.mu.RUnlock() - for _, sub := range s.Subscribers { - if sub.Ignored { - subs = append(subs, sub) + subs := make([]*Subscriber, 0, len(s.Subscribers)) + + for idx := range s.Subscribers { + if s.Subscribers[idx].Ignored { + subs = append(subs, s.Subscribers[idx]) } } diff --git a/subscriber_test.go b/subscriber_test.go index c7f741a..537afd6 100644 --- a/subscriber_test.go +++ b/subscriber_test.go @@ -1,64 +1,63 @@ package subscribe -/* XXX: a few new methods require tests. */ import ( "testing" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestCreateSub(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub := &Subscribe{Events: new(Events)} sub.CreateSub("myContacNameTest", "apiValueHere", true, false) - assert.EqualValues(1, len(sub.Subscribers), "there must be one subscriber") - assert.True(sub.Subscribers[0].Admin, "admin must be true") - assert.False(sub.Subscribers[0].Ignored, "ignore must be false") + assertions.Len(sub.Subscribers, 1, "there must be one subscriber") + assertions.True(sub.Subscribers[0].Admin, "admin must be true") + assertions.False(sub.Subscribers[0].Ignored, "ignore must be false") // Update values for existing contact. sub.CreateSub("myContacNameTest", "apiValueHere", false, true) - assert.EqualValues(1, len(sub.Subscribers), "there must still be one subscriber") - assert.False(sub.Subscribers[0].Admin, "admin must be changed to false") - assert.True(sub.Subscribers[0].Ignored, "ignore must be changed to true") - assert.True(sub.Subscribers[0].Ignored, "ignore must be changed to true") - assert.EqualValues(sub.Subscribers[0].Contact, "myContacNameTest", "contact value is incorrect") - assert.EqualValues(sub.Subscribers[0].API, "apiValueHere", "api value is incorrect") + assertions.Len(sub.Subscribers, 1, "there must still be one subscriber") + assertions.False(sub.Subscribers[0].Admin, "admin must be changed to false") + assertions.True(sub.Subscribers[0].Ignored, "ignore must be changed to true") + assertions.Equal("myContacNameTest", sub.Subscribers[0].Contact, "contact value is incorrect") + assertions.Equal("apiValueHere", sub.Subscribers[0].API, "api value is incorrect") // Add another contact. sub.CreateSub("myContacName2Test", "apiValueHere", false, true) - assert.EqualValues(2, len(sub.Subscribers), "there must be two subscribers") - assert.NotNil(sub.Subscribers[1].Events, "events map must not be nil") + assertions.Len(sub.Subscribers, 2, "there must be two subscribers") + assertions.NotNil(sub.Subscribers[1].Events, "events map must not be nil") } func TestGetSubscriber(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub := &Subscribe{Events: new(Events)} // Test missing subscriber _, err := sub.GetSubscriber("im not here", "fake") - assert.EqualValues(ErrSubscriberNotFound, err, "must have a subscriber not found error") + assertions.Equal(ErrSubscriberNotFound, err, "must have a subscriber not found error") // Test getting real subscriber sub.CreateSub("myContacNameTest", "apiValueHere", true, false) _, err = sub.GetSubscriber("myContacNameTest", "apiValueHere") - assert.Nil(err, "must not produce an error getting existing subscriber") + assertions.NoError(err, "must not produce an error getting existing subscriber") } func TestAdmin(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub := &Subscribe{Events: new(Events)} // Test missing subscriber subs := sub.GetAdmins() - assert.EqualValues(0, len(subs), "there must be zero admin since none were added") + assertions.Empty(subs, "there must be zero admin since none were added") // Test getting real subscriber sub.CreateSub("myContacNameTest", "apiValueHere", true, false) @@ -66,18 +65,19 @@ func TestAdmin(t *testing.T) { sub.CreateSub("myContacNameTest3", "apiValueHere", false, false) subs = sub.GetAdmins() - assert.EqualValues(1, len(subs), "there must be one admin") + assertions.Len(subs, 1, "there must be one admin") + assertions.Equal("myContacNameTest", subs[0].Contact) } func TestIgnore(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub := &Subscribe{Events: new(Events)} // Test missing subscriber subs := sub.GetIgnored() - assert.EqualValues(0, len(subs), "there must be zero ignored users since none were added") + assertions.Empty(subs, "there must be zero ignored users since none were added") // Test getting real subscriber sub.CreateSub("myContacNameTest", "apiValueHere", false, true) @@ -86,5 +86,45 @@ func TestIgnore(t *testing.T) { sub.CreateSub("myContacNameTest3", "apiValueHere", false, false) subs = sub.GetIgnored() - assert.EqualValues(1, len(subs), "there must be one ignored user") + assertions.Len(subs, 1, "there must be one ignored user") + assertions.Equal("myContacNameTest", subs[0].Contact) +} + +func TestCreateSubWithID(t *testing.T) { + t.Parallel() + + sub := &Subscribe{Events: new(Events)} + assert.Nil(t, sub.CreateSubWithID(0, "contact", "api", true, false)) + + first := sub.CreateSubWithID(10, "contact", "api", true, false) + require.NotNil(t, first) + assert.EqualValues(t, 10, first.ID) + assert.True(t, first.Admin) + assert.False(t, first.Ignored) + assert.Len(t, sub.Subscribers, 1) + + second := sub.CreateSubWithID(10, "contact-new", "api", false, true) + require.NotNil(t, second) + assert.Same(t, first, second) + assert.False(t, second.Admin) + assert.True(t, second.Ignored) + assert.Equal(t, "contact", second.Contact) + assert.Len(t, sub.Subscribers, 1) +} + +func TestGetSubscriberByID(t *testing.T) { + t.Parallel() + + sub := &Subscribe{Events: new(Events)} + _, err := sub.GetSubscriberByID(0, "api") + assert.Equal(t, ErrSubscriberNotFound, err) + + sub.CreateSubWithID(99, "contact", "api", true, false) + got, err := sub.GetSubscriberByID(99, "api") + require.NoError(t, err) + require.NotNil(t, got) + assert.EqualValues(t, 99, got.ID) + + _, err = sub.GetSubscriberByID(99, "api2") + assert.Equal(t, ErrSubscriberNotFound, err) } diff --git a/subscription.go b/subscription.go index 2b9f28a..0ef851e 100644 --- a/subscription.go +++ b/subscription.go @@ -20,10 +20,12 @@ func (s *Subscriber) Subscribe(event string) error { // Call this method when your event fires, collect the subscribers and send // them notifications in your app. Subscribers can be people. Or functions. func (s *Subscribe) GetSubscribers(eventName string) []*Subscriber { - var subscribers []*Subscriber + s.mu.RLock() + defer s.mu.RUnlock() + subscribers := make([]*Subscriber, 0, len(s.Subscribers)) for _, sub := range s.Subscribers { - if !sub.Ignored && s.checkAPI(sub.API) && !sub.Events.IsPaused(eventName) { + if !sub.Ignored && s.checkAPILocked(sub.API) && !sub.Events.IsPaused(eventName) { subscribers = append(subscribers, sub) } } @@ -33,6 +35,13 @@ func (s *Subscribe) GetSubscribers(eventName string) []*Subscriber { // checkAPI just looks for a string in a slice of strings with a twist. func (s *Subscribe) checkAPI(api string) bool { + s.mu.RLock() + defer s.mu.RUnlock() + + return s.checkAPILocked(api) +} + +func (s *Subscribe) checkAPILocked(api string) bool { if len(s.EnableAPIs) < 1 { return true } @@ -46,8 +55,11 @@ func (s *Subscribe) checkAPI(api string) bool { return false } -// EventRemove obliterates an event and all subsciptions for it. +// EventRemove obliterates an event and all subscriptions for it. func (s *Subscribe) EventRemove(event string) { + s.mu.RLock() + defer s.mu.RUnlock() + s.Events.Remove(event) for _, sub := range s.Subscribers { diff --git a/subscription_test.go b/subscription_test.go index be3d342..40f9240 100644 --- a/subscription_test.go +++ b/subscription_test.go @@ -6,147 +6,158 @@ import ( "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestCheckAPI(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) - sub := &Subscribe{Events: new(Events)} - assert.True(sub.checkAPI("test_string"), "an empty slice must always return true") + subscriber := &Subscribe{Events: new(Events)} + assertions.True(subscriber.checkAPI("test_string"), "an empty slice must always return true") - sub.EnableAPIs = []string{"event", "test_string"} - assert.True(sub.checkAPI("test_string://event"), "test_string is an allowed api prefix") + subscriber.EnableAPIs = []string{"event", "test_string"} + assertions.True(subscriber.checkAPI("test_string://event"), "test_string is an allowed api prefix") - sub.EnableAPIs = []string{"event", "any"} - assert.True(sub.checkAPI("test_string"), "any as a slice value must return true") + subscriber.EnableAPIs = []string{"event", "any"} + assertions.True(subscriber.checkAPI("test_string"), "any as an allowed value must return true") - sub.EnableAPIs = []string{"event", "all"} - assert.True(sub.checkAPI("test_string"), "all as a slice value must return true") + subscriber.EnableAPIs = []string{"event", "all"} + assertions.True(subscriber.checkAPI("test_string"), "all as an allowed value must return true") - sub.EnableAPIs = []string{"event", "test_string"} - assert.True(sub.checkAPI("test_string"), "test_string is an allowed api") + subscriber.EnableAPIs = []string{"event", "test_string"} + assertions.True(subscriber.checkAPI("test_string"), "test_string is an allowed api") - sub.EnableAPIs = []string{"event", "test_string2"} - assert.False(sub.checkAPI("test_string"), "test_string is not an allowed api") + subscriber.EnableAPIs = []string{"event", "test_string2"} + assertions.False(subscriber.checkAPI("test_string"), "test_string is not an allowed api") } func TestUnSubscribe(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub := &Subscribe{Events: new(Events)} - // Add 1 user and 3 subscriptions. - user := sub.CreateSub("myContacNameTest", "apiValueHere", true, true) - assert.Nil(user.Subscribe("event_name")) - assert.Nil(user.Subscribe("event_name2")) - assert.Nil(user.Subscribe("event_name3")) + // Add 1 subscriber and 3 subscriptions. + subscriber := sub.CreateSub("myContacNameTest", "apiValueHere", true, true) + require.NoError(t, subscriber.Subscribe("event_name")) + require.NoError(t, subscriber.Subscribe("event_name2")) + require.NoError(t, subscriber.Subscribe("event_name3")) // Make sure we can't add the same event twice. - assert.EqualValues(ErrEventExists, user.Subscribe("event_name3"), "duplicate event allowed") + assertions.Equal(ErrEventExists, subscriber.Subscribe("event_name3"), "duplicate event allowed") // Remove a subscription. - user.Events.Remove("event_name3") - assert.EqualValues(2, len(sub.Subscribers[0].Events.Map), "there must be two subscriptions remaining") + subscriber.Events.Remove("event_name3") + assertions.Len(sub.Subscribers[0].Events.Map, 2, "there must be two subscriptions remaining") // Remove another. - user.Events.Remove("event_name2") - assert.EqualValues(1, len(sub.Subscribers[0].Events.Map), "there must be one subscription remaining") - user.Events.Remove("event_name_not_here") + subscriber.Events.Remove("event_name2") + assertions.Len(sub.Subscribers[0].Events.Map, 1, "there must be one subscription remaining") + subscriber.Events.Remove("event_name_not_here") } func TestPause(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub := &Subscribe{Events: new(Events)} - user := sub.CreateSub("contact", "api", true, false) - assert.Nil(user.Subscribe("eventName")) + subscriber := sub.CreateSub("contact", "api", true, false) + require.NoError(t, subscriber.Subscribe("eventName")) // Make sure pausing a missing event returns the proper error. - assert.EqualValues(ErrEventNotFound, user.Events.Pause("fake event", 0)) + assertions.Equal(ErrEventNotFound, subscriber.Events.Pause("fake event", 0)) // Testing a real unpause. - assert.Nil(user.Events.Pause("eventName", 0)) - assert.WithinDuration(time.Now(), sub.Subscribers[0].Events.Map["eventName"].Pause, 1*time.Second) + require.NoError(t, subscriber.Events.Pause("eventName", 0)) + assertions.WithinDuration(time.Now(), sub.Subscribers[0].Events.Map["eventName"].Pause, 1*time.Second) // Testing a real pause. - assert.Nil(user.Events.Pause("eventName", 3600*time.Second)) - assert.WithinDuration(time.Now().Add(3600*time.Second), - sub.Subscribers[0].Events.Map["eventName"].Pause, 1*time.Second) + require.NoError(t, subscriber.Events.Pause("eventName", 3600*time.Second)) + assertions.WithinDuration( + time.Now().Add(3600*time.Second), + sub.Subscribers[0].Events.Map["eventName"].Pause, + 1*time.Second, + ) } func TestIsPaused(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub := &Subscribe{Events: new(Events)} - user := sub.CreateSub("contact", "api", true, false) + subscriber := sub.CreateSub("contact", "api", true, false) - // Go back and fourth a few times. - assert.Nil(user.Subscribe("eventName")) - assert.Nil(user.Events.Pause("eventName", 0)) - assert.False(user.Events.IsPaused("eventName")) - assert.Nil(user.Events.Pause("eventName", 10*time.Second)) - assert.True(user.Events.IsPaused("eventName")) - assert.Nil(user.Events.UnPause("eventName")) - assert.False(user.Events.IsPaused("eventName")) + // Go back and forth a few times. + require.NoError(t, subscriber.Subscribe("eventName")) + require.NoError(t, subscriber.Events.Pause("eventName", 0)) + assertions.False(subscriber.Events.IsPaused("eventName")) + require.NoError(t, subscriber.Events.Pause("eventName", 10*time.Second)) + assertions.True(subscriber.Events.IsPaused("eventName")) + require.NoError(t, subscriber.Events.UnPause("eventName")) + assertions.False(subscriber.Events.IsPaused("eventName")) // Missing event is always paused. - assert.True(user.Events.IsPaused("missingEvent")) + assertions.True(subscriber.Events.IsPaused("missingEvent")) } func TestSubscriptions(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub := &Subscribe{Events: new(Events)} - user := sub.CreateSub("contact", "api", true, false) + subscriber := sub.CreateSub("contact", "api", true, false) events := []string{"eventName", "eventName1", "eventName3", "eventName5"} sort.Strings(events) for _, e := range events { - assert.Nil(user.Subscribe(e)) + require.NoError(t, subscriber.Subscribe(e)) } - assert.Equal(events, user.Events.Names(), "wrong subscriptions provided") + assertions.Equal(events, subscriber.Events.Names(), "wrong subscriptions provided") } func TestGetSubscribers(t *testing.T) { t.Parallel() - assert := assert.New(t) + assertions := assert.New(t) sub := &Subscribe{Events: new(Events)} subs := sub.GetSubscribers("evn") - assert.EqualValues(0, len(subs), "there must be no subscribers") + assertions.Empty(subs, "there must be no subscribers") // Add 1 subscriber and 3 subscriptions. - user := sub.CreateSub("myContacNameTest", "apiValueHere", true, false) - assert.Nil(user.Subscribe("event_name")) - assert.Nil(user.Subscribe("event_name2")) - assert.Nil(user.Subscribe("event_name3")) + subscriber := sub.CreateSub("myContacNameTest", "apiValueHere", true, false) + require.NoError(t, subscriber.Subscribe("event_name")) + require.NoError(t, subscriber.Subscribe("event_name2")) + require.NoError(t, subscriber.Subscribe("event_name3")) // Add 1 more subscriber and 3 more subscriptions, 2 paused. - user = sub.CreateSub("myContacNameTest2", "apiValueHere", true, false) - assert.Nil(user.Subscribe("event_name")) - assert.Nil(user.Subscribe("event_name2")) - assert.Nil(user.Subscribe("event_name3")) - assert.Nil(user.Events.Pause("event_name2", 10*time.Second)) - assert.Nil(user.Events.Pause("event_name3", 10*time.Minute)) + subscriber = sub.CreateSub("myContacNameTest2", "apiValueHere", true, false) + require.NoError(t, subscriber.Subscribe("event_name")) + require.NoError(t, subscriber.Subscribe("event_name2")) + require.NoError(t, subscriber.Subscribe("event_name3")) + require.NoError(t, subscriber.Events.Pause("event_name2", 10*time.Second)) + require.NoError(t, subscriber.Events.Pause("event_name3", 10*time.Minute)) // Add another ignore subscriber with 1 subscription. - user = sub.CreateSub("myContacNameTest3", "apiValueHere", true, true) - assert.Nil(user.Subscribe("event_name")) + subscriber = sub.CreateSub("myContacNameTest3", "apiValueHere", true, true) + require.NoError(t, subscriber.Subscribe("event_name")) // Test that ignore keeps the ignored subscriber out. - assert.EqualValues(2, len(sub.GetSubscribers("event_name")), "there must be 2 subscribers") + subs = sub.GetSubscribers("event_name") + assertions.Len(subs, 2, "there must be 2 subscribers") + assertions.ElementsMatch([]string{"myContacNameTest", "myContacNameTest2"}, []string{subs[0].Contact, subs[1].Contact}) // Test that resume time keeps a subscriber out. - assert.EqualValues(1, len(sub.GetSubscribers("event_name2")), "there must be 1 subscriber") - assert.EqualValues(1, len(sub.GetSubscribers("event_name3")), "there must be 1 subscriber") + subs = sub.GetSubscribers("event_name2") + assertions.Len(subs, 1, "there must be 1 subscriber") + assertions.Equal("myContacNameTest", subs[0].Contact) + + subs = sub.GetSubscribers("event_name3") + assertions.Len(subs, 1, "there must be 1 subscriber") + assertions.Equal("myContacNameTest", subs[0].Contact) } diff --git a/types.go b/types.go index 08b0649..ac105c5 100644 --- a/types.go +++ b/types.go @@ -1,28 +1,28 @@ package subscribe import ( - "fmt" + "errors" "sync" "time" ) var ( // ErrSubscriberNotFound is returned any time a requested subscriber does not exist. - ErrSubscriberNotFound = fmt.Errorf("subscriber not found") + ErrSubscriberNotFound = errors.New("subscriber not found") // ErrEventNotFound is returned when a requested event has not been created. - ErrEventNotFound = fmt.Errorf("event not found") + ErrEventNotFound = errors.New("event not found") // ErrEventExists is returned when a new event with an existing name is created. - ErrEventExists = fmt.Errorf("event already exists") + ErrEventExists = errors.New("event already exists") ) // Rules contains the pause time and rules for a subscriber's event subscription. // Rules are unused by the library and available for consumers. type Rules struct { - Pause time.Time `json:"pause"` - D map[string]time.Duration - I map[string]int - S map[string]string - T map[string]time.Time + Pause time.Time `json:"pause"` + D map[string]time.Duration `json:"durations"` + I map[string]int `json:"integers"` + S map[string]string `json:"strings"` + T map[string]time.Time `json:"times"` } // Subscriber describes the contact info and subscriptions for a person. @@ -30,7 +30,7 @@ type Subscriber struct { // ID is optional. If it provided, this is used as the _match_. ID int64 `json:"id"` // Meta is optional. This library does not use this value. - Meta map[string]interface{} `json:"meta"` + Meta map[string]any `json:"meta"` // API is the type of API the subscriber is subscribed with. Used to filter results. API string `json:"api"` // Contact is the contact info used in the API to send the subscriber a notification. @@ -50,8 +50,8 @@ type Subscriber struct { type Events struct { // Map is the events/rules map. Use the provided methods to interact with it. Map map[string]*Rules `json:"eventsMap"` - // sync.RWMutex locks and unlocks the Events map - sync.RWMutex + // sync.mu locks and unlocks the Events map + mu sync.RWMutex } // Subscribe is the data needed to initialize this module. @@ -59,6 +59,8 @@ type Subscribe struct { // EnableAPIs sets the allowed APIs. Only subscriptions that have an API // with a prefix in this list will return from the GetSubscribers() method. EnableAPIs []string `json:"enabledApis"` // imessage, skype, pushover, email, slack, growl, all, any + // mu protects mutable Subscribe fields. + mu sync.RWMutex // stateFile is the db location, like: /usr/local/var/lib/motifini/subscribers.json stateFile string // Events stores a list of arbitrary events. Use the included methods to interact with it.