diff --git a/cmd/coscout/commands/daemon.go b/cmd/coscout/commands/daemon.go index ca91006..9b5f9a3 100644 --- a/cmd/coscout/commands/daemon.go +++ b/cmd/coscout/commands/daemon.go @@ -15,8 +15,10 @@ package commands import ( + "context" "os" "os/signal" + "sync" "sync/atomic" "syscall" @@ -31,7 +33,56 @@ import ( ) type authState struct { - isAuthed atomic.Bool + lifecycle atomic.Int32 + cancelMu sync.Mutex + cancel context.CancelFunc +} + +const ( + daemonIdle int32 = iota + daemonRunning + daemonStopping +) + +func (s *authState) tryStart() (context.Context, bool) { + ctx, cancel := context.WithCancel(context.Background()) + + s.cancelMu.Lock() + defer s.cancelMu.Unlock() + if !s.lifecycle.CompareAndSwap(daemonIdle, daemonRunning) { + cancel() + return nil, false + } + s.cancel = cancel + return ctx, true +} + +func (s *authState) requestStop() bool { + if !s.lifecycle.CompareAndSwap(daemonRunning, daemonStopping) { + return false + } + + s.cancelMu.Lock() + cancel := s.cancel + s.cancelMu.Unlock() + if cancel != nil { + cancel() + } + return true +} + +func (s *authState) daemonStopped() { + s.cancelMu.Lock() + cancelRunAndMarkIdle(s.cancel, &s.lifecycle) + s.cancel = nil + s.cancelMu.Unlock() +} + +func cancelRunAndMarkIdle(cancel context.CancelFunc, lifecycle *atomic.Int32) { + if cancel != nil { + cancel() + } + lifecycle.Store(daemonIdle) } func NewDaemonCommand(cfgPath *string) *cobra.Command { @@ -42,7 +93,10 @@ func NewDaemonCommand(cfgPath *string) *cobra.Command { storageDB := storage.NewBoltDB(config.GetDBPath()) confManager := config.InitConfManager(*cfgPath, &storageDB) - appConf := confManager.LoadOnce() + appConf, err := confManager.LoadStartup() + if err != nil { + log.Fatalf("Unable to load config file from %s: %v", *cfgPath, err) + } log.Infof("Load config file from %s", *cfgPath) registerChan := make(chan model.DeviceStatusResponse, 10) @@ -66,28 +120,40 @@ func NewDaemonCommand(cfgPath *string) *cobra.Command { } func run(confManager *config.ConfManager, reqClient *api.RequestClient, registerChan chan model.DeviceStatusResponse) { - startChan := make(chan bool, 1) - exitChan := make(chan bool, 1) errorChan := make(chan error, 100) state := &authState{} - state.isAuthed.Store(false) for deviceStatus := range registerChan { if deviceStatus.Authorized { - log.Info("Device is authorized. Performing actions...") + startDaemon(state, confManager, reqClient, errorChan) + continue + } - if !state.isAuthed.Load() { - go daemon.Run(confManager, reqClient, startChan, exitChan, errorChan) - startChan <- true - } - state.isAuthed.Store(true) - } else { - log.Warn("Device is not authorized, waiting...") + stopDaemon(state) + } +} - if state.isAuthed.Load() { - exitChan <- true - } - state.isAuthed.Store(false) - } +func startDaemon( + state *authState, + confManager *config.ConfManager, + reqClient *api.RequestClient, + errorChan chan error, +) { + log.Info("Device is authorized. Performing actions...") + ctx, started := state.tryStart() + if !started { + return } + + go func() { + defer state.daemonStopped() + if err := daemon.Run(ctx, confManager, reqClient, errorChan); err != nil { + log.Errorf("Daemon stopped: %v", err) + } + }() +} + +func stopDaemon(state *authState) { + log.Warn("Device is not authorized, waiting...") + state.requestStop() } diff --git a/cmd/coscout/commands/daemon_test.go b/cmd/coscout/commands/daemon_test.go new file mode 100644 index 0000000..8981a5d --- /dev/null +++ b/cmd/coscout/commands/daemon_test.go @@ -0,0 +1,93 @@ +// Copyright 2025 coScene +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package commands + +import ( + "context" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAuthStateStopFromFailedRunDoesNotCancelRestart(t *testing.T) { + t.Parallel() + + state := &authState{} + + firstCtx, started := state.tryStart() + require.True(t, started) + _, started = state.tryStart() + require.False(t, started) + require.True(t, state.requestStop()) + require.False(t, state.requestStop()) + require.ErrorIs(t, firstCtx.Err(), context.Canceled) + _, started = state.tryStart() + require.False(t, started) + + state.daemonStopped() + secondCtx, started := state.tryStart() + require.True(t, started) + require.NoError(t, secondCtx.Err(), "a previous stop must not cancel a restarted daemon") +} + +func TestAuthStateRepeatedAuthorizationTransitionsAreIdempotent(t *testing.T) { + t.Parallel() + + state := &authState{} + for range 10 { + runCtx, started := state.tryStart() + require.True(t, started) + + _, started = state.tryStart() + require.False(t, started, "duplicate authorization must not start another daemon") + require.True(t, state.requestStop()) + require.False(t, state.requestStop(), "duplicate unauthorization must not request another stop") + require.ErrorIs(t, runCtx.Err(), context.Canceled) + + state.daemonStopped() + } +} + +func TestDaemonStoppedCancelsOldRunBeforeRestart(t *testing.T) { + t.Parallel() + + state := &authState{} + firstCtx, started := state.tryStart() + require.True(t, started) + + state.daemonStopped() + + secondCtx, started := state.tryStart() + require.True(t, started) + require.ErrorIs(t, firstCtx.Err(), context.Canceled) + require.NoError(t, secondCtx.Err()) +} + +func TestCancelRunAndMarkIdleOrdersCancellationFirst(t *testing.T) { + t.Parallel() + + var lifecycle atomic.Int32 + lifecycle.Store(daemonRunning) + runCtx, cancel := context.WithCancel(t.Context()) + + cancelRunAndMarkIdle(func() { + require.Equal(t, daemonRunning, lifecycle.Load(), "idle became visible before the old run was canceled") + cancel() + }, &lifecycle) + + require.ErrorIs(t, runCtx.Err(), context.Canceled) + require.Equal(t, daemonIdle, lifecycle.Load()) +} diff --git a/internal/collector/collector.go b/internal/collector/collector.go index eb282a3..2f8bcf7 100644 --- a/internal/collector/collector.go +++ b/internal/collector/collector.go @@ -117,7 +117,14 @@ func Collect(ctx context.Context, reqClient *api.RequestClient, confManager *con } }() - appConfig := confManager.LoadWithRemote() + appConfig, configErr := confManager.LoadWithRemote() + if configErr != nil { + if appConfig == nil { + log.Errorf("Unable to load collector config: %v", configErr) + return + } + log.Warnf("Unable to reload collector config, using last-known-good config: %v", configErr) + } getStorage := confManager.GetStorage() err := handleRecordCaches(uploadChan, reqClient, appConfig, getStorage, recordSet) diff --git a/internal/collector/upload.go b/internal/collector/upload.go index fe7b4a5..bda540e 100644 --- a/internal/collector/upload.go +++ b/internal/collector/upload.go @@ -142,7 +142,14 @@ func Upload(ctx context.Context, reqClient *api.RequestClient, confManager *conf return } - appConfig := confManager.LoadWithRemote() + appConfig, configErr := confManager.LoadWithRemote() + if configErr != nil { + if appConfig == nil { + log.Errorf("Unable to load upload config: %v", configErr) + return + } + log.Warnf("Unable to reload upload config, using last-known-good config: %v", configErr) + } if appConfig != nil { enabled := appConfig.Upload.NetworkRule.Enabled blackInterfaces := appConfig.Upload.NetworkRule.BlackInterfaces @@ -219,7 +226,13 @@ func uploadFiles(ctx context.Context, reqClient *api.RequestClient, confManager } allCompleted := true - appConfig := confManager.LoadWithRemote() + appConfig, configErr := confManager.LoadWithRemote() + if configErr != nil { + if appConfig == nil { + return errors.Wrap(configErr, "load upload config") + } + log.Warnf("Unable to reload upload config, using last-known-good config: %v", configErr) + } getStorage := confManager.GetStorage() recordCache, err := recordCache.Reload() diff --git a/internal/config/manager.go b/internal/config/manager.go index 4b0d40c..5252929 100644 --- a/internal/config/manager.go +++ b/internal/config/manager.go @@ -16,6 +16,7 @@ package config import ( "strings" + "sync" "github.com/coscene-io/coscout/internal/storage" "github.com/coscene-io/coscout/pkg/constant" @@ -25,6 +26,7 @@ import ( "github.com/knadh/koanf/providers/file" "github.com/knadh/koanf/providers/rawbytes" "github.com/knadh/koanf/v2" + "github.com/pkg/errors" log "github.com/sirupsen/logrus" ) @@ -36,12 +38,19 @@ const ( type ConfManager struct { cfg string storage *storage.Storage + cache *configCache +} + +type configCache struct { + mu sync.RWMutex + lastGood *AppConfig } func InitConfManager(cfg string, s *storage.Storage) *ConfManager { return &ConfManager{ cfg: cfg, storage: s, + cache: &configCache{}, } } @@ -84,11 +93,11 @@ func (c ConfManager) getDefaultConfig() AppConfig { } } -func (c ConfManager) LoadOnce() AppConfig { +func (c ConfManager) LoadOnce() (AppConfig, error) { appConf := c.getDefaultConfig() if err := utils.ParseYAML(c.cfg, &appConf); err != nil { - log.Fatalf("unable to load cos config: %v", err) + return AppConfig{}, errors.Wrap(err, "unable to load cos config") } for _, f := range appConf.Import { if strings.HasPrefix(f, RemoteFilePrefix) { @@ -97,58 +106,99 @@ func (c ConfManager) LoadOnce() AppConfig { localPath := strings.TrimPrefix(f, LocalFilePrefix) if !utils.CheckReadPath(localPath) { - log.Warnf("local config file %s not exist", localPath) - continue + return AppConfig{}, errors.Errorf("unable to load local config %s: file does not exist or is not readable", localPath) } if err := utils.ParseYAML(localPath, &appConf); err != nil { - log.Fatalf("unable to load cos config: %v", err) + return AppConfig{}, errors.Wrapf(err, "unable to load local config %s", localPath) } } - return appConf + return appConf, nil +} + +// LoadStartup loads the local startup configuration and seeds it as an initial +// fallback. Remote imports are intentionally not required at this stage; the +// first successful LoadWithRemote replaces this baseline with the complete +// merged configuration. +func (c ConfManager) LoadStartup() (AppConfig, error) { + appConf, err := c.LoadOnce() + if err != nil { + return AppConfig{}, err + } + + c.setLastKnownGood(appConf) + return appConf, nil } -func (c ConfManager) LoadWithRemote() *AppConfig { - appConf := c.LoadOnce() +func (c ConfManager) LoadWithRemote() (*AppConfig, error) { + appConf, err := c.LoadOnce() + if err != nil { + return c.lastKnownGood(), err + } k := koanf.New(".") // Create a new koanf instance. for _, f := range appConf.Import { //nolint: nestif // no need to nest if if strings.HasPrefix(f, RemoteFilePrefix) { name := strings.TrimPrefix(f, RemoteFilePrefix) + if c.storage == nil || *c.storage == nil { + return c.lastKnownGood(), errors.Errorf("unable to load remote config %s: storage is not configured", name) + } remoteCache, err := (*c.storage).Get([]byte(constant.DeviceRemoteConfigBucket), []byte(name)) if err != nil { - log.Errorf("unable to get remote config: %v", err) - continue + return c.lastKnownGood(), errors.Wrapf(err, "unable to get remote config %s", name) } //nolint: gosimple // no need to simplify if remoteCache == nil || len(remoteCache) == 0 { - log.Errorf("remote config is empty") - continue + return c.lastKnownGood(), errors.Errorf("unable to load remote config %s: config is empty", name) } if err := k.Load(rawbytes.Provider(remoteCache), json.Parser()); err != nil { - log.Errorf("unable to load remote config: %v", err) - continue + return c.lastKnownGood(), errors.Wrapf(err, "unable to load remote config %s", name) } } else { localPath := strings.TrimPrefix(f, LocalFilePrefix) if !utils.CheckReadPath(localPath) { - log.Warnf("file %s not exist or has no permission", localPath) - continue + return c.lastKnownGood(), errors.Errorf("unable to load local config %s: file does not exist or is not readable", localPath) } if err := k.Load(file.Provider(localPath), yaml.Parser()); err != nil { - log.Errorf("unable to load local config: %v", err) - continue + return c.lastKnownGood(), errors.Wrapf(err, "unable to load local config %s", localPath) } } } - err := k.Unmarshal("", &appConf) - if err != nil { - log.Errorf("unable to unmarshal koanf: %v", err) + if err := k.Unmarshal("", &appConf); err != nil { + return c.lastKnownGood(), errors.Wrap(err, "unable to unmarshal config") + } + + c.setLastKnownGood(appConf) + return &appConf, nil +} + +func (c ConfManager) setLastKnownGood(appConf AppConfig) { + if c.cache == nil { + return + } + + c.cache.mu.Lock() + defer c.cache.mu.Unlock() + + c.cache.lastGood = &appConf +} + +func (c ConfManager) lastKnownGood() *AppConfig { + if c.cache == nil { + return nil + } + + c.cache.mu.RLock() + defer c.cache.mu.RUnlock() + + if c.cache.lastGood == nil { + return nil } + appConf := *c.cache.lastGood return &appConf } diff --git a/internal/config/manager_test.go b/internal/config/manager_test.go index 87fa502..15cb131 100644 --- a/internal/config/manager_test.go +++ b/internal/config/manager_test.go @@ -15,11 +15,295 @@ package config import ( + "errors" + "fmt" + "os" + "path/filepath" + "sync" "testing" + "github.com/coscene-io/coscout/internal/storage" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) +var ( + errReloadUnexpectedlySucceeded = errors.New("reload unexpectedly succeeded") + errReloadMissingLastKnownGood = errors.New("reload returned no last-known-good config") + errReloadPortMismatch = errors.New("reload port mismatch") + errStorageUnavailable = errors.New("storage unavailable") +) + +type configTestStorage struct { + mu sync.RWMutex + value []byte + getErr error +} + +func (s *configTestStorage) Put(_, _, value []byte) error { + s.mu.Lock() + defer s.mu.Unlock() + + s.value = append(s.value[:0], value...) + return nil +} + +func (s *configTestStorage) Get(_, _ []byte) ([]byte, error) { + s.mu.RLock() + defer s.mu.RUnlock() + + return append([]byte(nil), s.value...), s.getErr +} + +func (s *configTestStorage) Delete(_, _ []byte) error { + s.mu.Lock() + defer s.mu.Unlock() + + s.value = nil + return nil +} + +func (s *configTestStorage) Close() error { + return nil +} + +func (s *configTestStorage) Iter(_ []byte, _ func(key, value []byte) error) error { + return nil +} + +func (s *configTestStorage) setValue(value []byte) { + s.mu.Lock() + defer s.mu.Unlock() + + s.value = append(s.value[:0], value...) +} + +func (s *configTestStorage) setGetError(err error) { + s.mu.Lock() + defer s.mu.Unlock() + + s.getErr = err +} + +func TestLoadOnceReturnsInvalidYAMLError(t *testing.T) { + t.Parallel() + + configPath := filepath.Join(t.TempDir(), "cos.yaml") + require.NoError(t, os.WriteFile(configPath, []byte("http_server: ["), 0o600)) + + manager := InitConfManager(configPath, nil) + _, err := manager.LoadOnce() + + require.Error(t, err) +} + +func TestLoadWithRemoteReturnsLastKnownGoodOnReloadFailure(t *testing.T) { + t.Parallel() + + configPath := filepath.Join(t.TempDir(), "cos.yaml") + require.NoError(t, os.WriteFile(configPath, []byte("http_server:\n port: 12345\n"), 0o600)) + + manager := InitConfManager(configPath, nil) + first, err := manager.LoadWithRemote() + require.NoError(t, err) + require.NotNil(t, first) + require.Equal(t, 12345, first.HttpServer.Port) + + require.NoError(t, os.WriteFile(configPath, []byte("http_server: ["), 0o600)) + + reloaded, err := manager.LoadWithRemote() + require.Error(t, err) + require.NotNil(t, reloaded) + require.Equal(t, *first, *reloaded) +} + +func TestLoadWithRemoteLastKnownGoodIsConcurrentSafe(t *testing.T) { + t.Parallel() + + configPath := filepath.Join(t.TempDir(), "cos.yaml") + require.NoError(t, os.WriteFile(configPath, []byte("http_server:\n port: 12345\n"), 0o600)) + + manager := InitConfManager(configPath, nil) + loaded, err := manager.LoadWithRemote() + require.NoError(t, err) + require.NotNil(t, loaded) + + require.NoError(t, os.WriteFile(configPath, []byte("http_server: ["), 0o600)) + + managerCopy := *manager + const readers = 20 + var wg sync.WaitGroup + results := make(chan error, readers) + wg.Add(readers) + for i := range readers { + go func(manager ConfManager) { + defer wg.Done() + reloaded, loadErr := manager.LoadWithRemote() + if loadErr == nil { + results <- errReloadUnexpectedlySucceeded + return + } + if reloaded == nil { + results <- errReloadMissingLastKnownGood + return + } + if reloaded.HttpServer.Port != 12345 { + results <- fmt.Errorf("%w: got %d, want 12345", errReloadPortMismatch, reloaded.HttpServer.Port) + return + } + results <- nil + }([]ConfManager{*manager, managerCopy}[i%2]) + } + wg.Wait() + close(results) + for result := range results { + require.NoError(t, result) + } +} + +func TestLoadWithRemoteRejectsInvalidImportedConfig(t *testing.T) { + t.Parallel() + + configDir := t.TempDir() + importPath := filepath.Join(configDir, "import.yaml") + configPath := filepath.Join(configDir, "cos.yaml") + require.NoError(t, os.WriteFile(importPath, []byte("http_server:\n port: 23456\n"), 0o600)) + require.NoError(t, os.WriteFile( + configPath, + []byte(fmt.Sprintf("__import__:\n - file://%s\n", importPath)), + 0o600, + )) + + manager := InitConfManager(configPath, nil) + loaded, err := manager.LoadWithRemote() + require.NoError(t, err) + require.Equal(t, 23456, loaded.HttpServer.Port) + + require.NoError(t, os.WriteFile(importPath, []byte("http_server: ["), 0o600)) + + reloaded, err := manager.LoadWithRemote() + require.Error(t, err) + require.NotNil(t, reloaded) + require.Equal(t, 23456, reloaded.HttpServer.Port) +} + +func TestLoadWithRemoteRejectsInvalidRemoteConfig(t *testing.T) { + t.Parallel() + + configPath := filepath.Join(t.TempDir(), "cos.yaml") + require.NoError(t, os.WriteFile(configPath, []byte("__import__:\n - cos://remote\n"), 0o600)) + + backend := &configTestStorage{value: []byte(`{"http_server":{"port":34567}}`)} + var store storage.Storage = backend + manager := InitConfManager(configPath, &store) + + loaded, err := manager.LoadWithRemote() + require.NoError(t, err) + require.Equal(t, 34567, loaded.HttpServer.Port) + + backend.setValue([]byte("{")) + + reloaded, err := manager.LoadWithRemote() + require.Error(t, err) + require.NotNil(t, reloaded) + require.Equal(t, 34567, reloaded.HttpServer.Port) +} + +func TestLoadWithRemoteDoesNotReplaceLastKnownGoodWhenRemoteIsUnavailable(t *testing.T) { + t.Parallel() + + configPath := filepath.Join(t.TempDir(), "cos.yaml") + require.NoError(t, os.WriteFile(configPath, []byte("__import__:\n - cos://remote\n"), 0o600)) + + tests := []struct { + name string + update func(*configTestStorage) + }{ + { + name: "read error", + update: func(backend *configTestStorage) { + backend.setGetError(errStorageUnavailable) + }, + }, + { + name: "empty value", + update: func(backend *configTestStorage) { + backend.setValue(nil) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + backend := &configTestStorage{value: []byte(`{"http_server":{"port":34567}}`)} + var store storage.Storage = backend + manager := InitConfManager(configPath, &store) + + loaded, err := manager.LoadWithRemote() + require.NoError(t, err) + require.Equal(t, 34567, loaded.HttpServer.Port) + + tt.update(backend) + + reloaded, err := manager.LoadWithRemote() + require.Error(t, err) + require.NotNil(t, reloaded) + require.Equal(t, 34567, reloaded.HttpServer.Port) + }) + } +} + +func TestLoadWithRemoteDoesNotReplaceLastKnownGoodWhenLocalImportIsMissing(t *testing.T) { + t.Parallel() + + configDir := t.TempDir() + importPath := filepath.Join(configDir, "import.yaml") + configPath := filepath.Join(configDir, "cos.yaml") + require.NoError(t, os.WriteFile(importPath, []byte("http_server:\n port: 23456\n"), 0o600)) + require.NoError(t, os.WriteFile( + configPath, + []byte(fmt.Sprintf("__import__:\n - file://%s\n", importPath)), + 0o600, + )) + + manager := InitConfManager(configPath, nil) + loaded, err := manager.LoadWithRemote() + require.NoError(t, err) + require.Equal(t, 23456, loaded.HttpServer.Port) + + require.NoError(t, os.Remove(importPath)) + + reloaded, err := manager.LoadWithRemote() + require.Error(t, err) + require.NotNil(t, reloaded) + require.Equal(t, 23456, reloaded.HttpServer.Port) +} + +func TestLoadStartupSeedsFallbackBeforeRemoteConfigIsAvailable(t *testing.T) { + t.Parallel() + + configPath := filepath.Join(t.TempDir(), "cos.yaml") + require.NoError(t, os.WriteFile( + configPath, + []byte("http_server:\n port: 45678\n__import__:\n - cos://remote\n"), + 0o600, + )) + + manager := InitConfManager(configPath, nil) + startup, err := manager.LoadStartup() + require.NoError(t, err) + require.Equal(t, 45678, startup.HttpServer.Port) + + require.NoError(t, os.WriteFile(configPath, []byte("http_server: ["), 0o600)) + + reloaded, err := manager.LoadWithRemote() + require.Error(t, err) + require.NotNil(t, reloaded) + require.Equal(t, 45678, reloaded.HttpServer.Port) +} + func TestDefaultHTTPQueuedBytesLimit(t *testing.T) { t.Parallel() diff --git a/internal/daemon/daemon.go b/internal/daemon/daemon.go index b99fb8b..609db2e 100644 --- a/internal/daemon/daemon.go +++ b/internal/daemon/daemon.go @@ -33,13 +33,26 @@ import ( "google.golang.org/protobuf/encoding/protojson" ) -func Run(confManager *config.ConfManager, reqClient *api.RequestClient, startChan chan bool, finishChan chan bool, errorChan chan error) { - <-startChan - - ctx, cancel := context.WithCancel(context.Background()) +func Run( + ctx context.Context, + confManager *config.ConfManager, + reqClient *api.RequestClient, + errorChan chan error, +) error { + runCtx, cancel := context.WithCancel(ctx) defer cancel() - appConfig := confManager.LoadWithRemote() + if runCtx.Err() != nil { + return nil + } + + appConfig, err := confManager.LoadWithRemote() + if err != nil { + if appConfig == nil { + return err + } + log.Warnf("Unable to reload config, continuing with last-known-good config: %v", err) + } queuedBytesLimit := appConfig.HttpServer.QueuedBytesLimit if queuedBytesLimit <= 0 { log.Warnf( @@ -77,20 +90,35 @@ func Run(confManager *config.ConfManager, reqClient *api.RequestClient, startCha var masterServer *master.Server var masterConfig *config.MasterConfig var fileManager *master.FileManager + var masterResult <-chan error if appConfig.MasterSlave.Enabled { masterConfig = config.DefaultMasterConfig() masterServer = master.NewServer(masterConfig.Port, masterConfig) + result := make(chan error, 1) + masterResult = result go func() { - if err := masterServer.Start(ctx); err != nil { + result <- masterServer.Start(runCtx) + }() + + select { + case <-masterServer.Ready(): + if runCtx.Err() != nil { + err := <-masterResult + cancel() + return err + } + log.Infof("Master server started on port %d", masterConfig.Port) + case err := <-masterResult: + if err != nil { log.Errorf("Master server failed: %v", err) - select { - case errorChan <- err: - default: - log.Warnf("Error channel is full, dropping error: %v", err) - } } - }() - log.Infof("Master server started on port %d", masterConfig.Port) + cancel() + return err + case <-runCtx.Done(): + err := <-masterResult + cancel() + return err + } // Set up FileManager for slave file handling masterClient := master.NewClient(masterConfig) @@ -100,7 +128,7 @@ func Run(confManager *config.ConfManager, reqClient *api.RequestClient, startCha // Start collector go func() { - err := collector.Collect(ctx, reqClient, confManager, fileManager, pubSub, errorChan) + err := collector.Collect(runCtx, reqClient, confManager, fileManager, pubSub, errorChan) if err != nil { select { case errorChan <- err: @@ -120,7 +148,7 @@ func Run(confManager *config.ConfManager, reqClient *api.RequestClient, startCha ticker.Reset(config.RefreshRemoteConfigInterval) //nolint: contextcheck // context is checked in the parent goroutine refreshRemoteConfig(confManager, reqClient) - case <-ctx.Done(): + case <-runCtx.Done(): log.Infof("Daemon ticker goroutine done") return } @@ -128,7 +156,7 @@ func Run(confManager *config.ConfManager, reqClient *api.RequestClient, startCha }(ticker) // Start core services - go core.SendHeartbeat(ctx, reqClient, confManager.GetStorage(), errorChan) + go core.SendHeartbeat(runCtx, reqClient, confManager.GetStorage(), errorChan) // Start task handler with optional master-slave enhancement taskHandler := mod.NewModHandler(*reqClient, *confManager, pubSub, errorChan, constant.TaskModType, ruleQueuedBytes) @@ -140,7 +168,7 @@ func Run(confManager *config.ConfManager, reqClient *api.RequestClient, startCha log.Info("Task handler enhanced with master-slave support") } } - go taskHandler.Run(ctx) + go taskHandler.Run(runCtx) // Start rule handler with optional master-slave enhancement ruleHandler := mod.NewModHandler(*reqClient, *confManager, pubSub, errorChan, constant.RuleModType, ruleQueuedBytes) @@ -152,17 +180,46 @@ func Run(confManager *config.ConfManager, reqClient *api.RequestClient, startCha log.Info("Rule handler enhanced with master-slave support") } } - go ruleHandler.Run(ctx) + go ruleHandler.Run(runCtx) // Start HTTP handler - go mod.NewModHandler(*reqClient, *confManager, pubSub, errorChan, constant.HttpModType, ruleQueuedBytes).Run(ctx) + go mod.NewModHandler(*reqClient, *confManager, pubSub, errorChan, constant.HttpModType, ruleQueuedBytes).Run(runCtx) log.Info("Daemon started successfully") - <-finishChan + if masterResult != nil { + return waitForMasterAndCancel(runCtx, masterResult, cancel) + } + + <-runCtx.Done() + cancel() + return nil +} + +// waitForMasterAndCancel returns a runtime master failure immediately and +// cancels all other daemon services before returning. On daemon cancellation +// it still waits for Server.Start to finish shutdown so callers cannot start a +// replacement daemon before the listening port is released. +func waitForMasterAndCancel( + ctx context.Context, + masterResult <-chan error, + cancel context.CancelFunc, +) error { + var err error + select { + case err = <-masterResult: + case <-ctx.Done(): + err = <-masterResult + } + cancel() + return err } func refreshRemoteConfig(confManager *config.ConfManager, reqClient *api.RequestClient) { - appConfig := confManager.LoadOnce() + appConfig, err := confManager.LoadOnce() + if err != nil { + log.Errorf("Unable to load local config while refreshing remote config: %v", err) + return + } if len(appConfig.Import) == 0 { return } diff --git a/internal/daemon/daemon_test.go b/internal/daemon/daemon_test.go new file mode 100644 index 0000000..b62f60c --- /dev/null +++ b/internal/daemon/daemon_test.go @@ -0,0 +1,106 @@ +// Copyright 2025 coScene +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package daemon + +import ( + "context" + "errors" + "net" + "os" + "path/filepath" + "testing" + "time" + + "github.com/coscene-io/coscout/internal/config" + "github.com/stretchr/testify/require" +) + +var errTestMasterServeFailed = errors.New("master serve failed") + +func TestRunReturnsMasterBindError(t *testing.T) { + t.Parallel() + + // #nosec G102 -- this test must reserve every interface used by the daemon. + listener, err := net.Listen("tcp", ":22525") + if err == nil { + t.Cleanup(func() { + require.NoError(t, listener.Close()) + }) + } + + configPath := filepath.Join(t.TempDir(), "cos.yaml") + require.NoError(t, os.WriteFile(configPath, []byte(` +master_slave: + enabled: true +`), 0o600)) + + confManager := config.InitConfManager(configPath, nil) + errorChan := make(chan error, 1) + runErr := make(chan error, 1) + + go func() { + runErr <- Run(t.Context(), confManager, nil, errorChan) + }() + + select { + case err := <-runErr: + require.Error(t, err) + case <-time.After(time.Second): + t.Fatal("Run did not return the master bind error") + } + + select { + case err := <-errorChan: + t.Fatalf("startup error was also sent as a non-fatal runtime error: %v", err) + default: + } +} + +func TestWaitForMasterReturnsRuntimeError(t *testing.T) { + t.Parallel() + + masterResult := make(chan error, 1) + masterResult <- errTestMasterServeFailed + runCtx, cancel := context.WithCancel(t.Context()) + + require.ErrorIs(t, waitForMasterAndCancel(runCtx, masterResult, cancel), errTestMasterServeFailed) + require.ErrorIs(t, runCtx.Err(), context.Canceled) +} + +func TestWaitForMasterCancellationWaitsForShutdown(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(t.Context()) + masterResult := make(chan error) + waitResult := make(chan error, 1) + go func() { + waitResult <- waitForMasterAndCancel(ctx, masterResult, cancel) + }() + + cancel() + select { + case <-waitResult: + t.Fatal("waitForMaster returned before master shutdown completed") + case <-time.After(20 * time.Millisecond): + } + + masterResult <- nil + select { + case err := <-waitResult: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("waitForMaster did not return after master shutdown completed") + } +} diff --git a/internal/master/server.go b/internal/master/server.go index 04d5e5a..76fe4f3 100644 --- a/internal/master/server.go +++ b/internal/master/server.go @@ -22,18 +22,24 @@ import ( "net" "net/http" "strings" + "sync" "time" "github.com/coscene-io/coscout/internal/config" log "github.com/sirupsen/logrus" ) +var errServerAlreadyStarted = errors.New("master server already started") + // Server master server. type Server struct { registry *SlaveRegistry config *config.MasterConfig server *http.Server port int + ready chan struct{} + startMu sync.Mutex + started bool } // NewServer creates a new master server. @@ -54,6 +60,7 @@ func NewServer(port int, masterConfig *config.MasterConfig) *Server { config: masterConfig, server: server, port: port, + ready: make(chan struct{}), } // Register routes @@ -66,26 +73,64 @@ func NewServer(port int, masterConfig *config.MasterConfig) *Server { return s } +// Ready is closed once the server has successfully bound its listening port. +func (s *Server) Ready() <-chan struct{} { + return s.ready +} + // Start starts the server. func (s *Server) Start(ctx context.Context) error { - // Start cleanup goroutine - go s.cleanupRoutine(ctx) + s.startMu.Lock() + if s.started { + s.startMu.Unlock() + return errServerAlreadyStarted + } + s.started = true + s.startMu.Unlock() log.Infof("Master server starting on port %d", s.port) + listener, err := net.Listen("tcp", s.server.Addr) + if err != nil { + return fmt.Errorf("listen on %s: %w", s.server.Addr, err) + } + close(s.ready) + + serverCtx, cancel := context.WithCancel(ctx) + defer cancel() + + // Start cleanup only after the listening socket has been acquired. + go s.cleanupRoutine(serverCtx) + + serveErr := make(chan error, 1) go func() { - if err := s.server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) { - log.Errorf("Master server failed: %v", err) - } + serveErr <- s.server.Serve(listener) }() - <-ctx.Done() - log.Info("Master server shutting down...") + select { + case err := <-serveErr: + if err == nil || errors.Is(err, http.ErrServerClosed) { + return nil + } + return fmt.Errorf("serve master requests: %w", err) + case <-ctx.Done(): + log.Info("Master server shutting down...") - shutdownCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) - defer cancel() + shutdownCtx, shutdownCancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer shutdownCancel() - return s.server.Shutdown(shutdownCtx) + if err := s.server.Shutdown(shutdownCtx); err != nil { + _ = s.server.Close() + <-serveErr + return fmt.Errorf("shutdown master server: %w", err) + } + + err := <-serveErr + if err != nil && !errors.Is(err, http.ErrServerClosed) { + return fmt.Errorf("serve master requests: %w", err) + } + return nil + } } // GetRegistry returns the slave registry. diff --git a/internal/master/server_test.go b/internal/master/server_test.go new file mode 100644 index 0000000..ee7a366 --- /dev/null +++ b/internal/master/server_test.go @@ -0,0 +1,104 @@ +// Copyright 2025 coScene +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package master + +import ( + "context" + "fmt" + "net" + "testing" + "time" + + "github.com/coscene-io/coscout/internal/config" + "github.com/stretchr/testify/require" +) + +func TestServerStartReturnsBindErrorWithoutReportingReady(t *testing.T) { + t.Parallel() + + // #nosec G102 -- this test must reserve every interface used by the server. + listener, err := net.Listen("tcp", ":0") + require.NoError(t, err) + t.Cleanup(func() { + require.NoError(t, listener.Close()) + }) + + addr, ok := listener.Addr().(*net.TCPAddr) + require.True(t, ok) + port := addr.Port + server := NewServer(port, config.DefaultMasterConfig()) + + startErr := make(chan error, 1) + go func() { + startErr <- server.Start(t.Context()) + }() + + select { + case err := <-startErr: + require.Error(t, err) + case <-time.After(time.Second): + t.Fatal("Start did not return the bind error") + } + + select { + case <-server.Ready(): + t.Fatal("server reported ready despite the bind error") + default: + } +} + +func TestServerStartReportsReadyAfterListenAndStopsOnCancellation(t *testing.T) { + t.Parallel() + + // Reserve a concrete port first so the post-shutdown bind verifies that the + // exact listener owned by this server has been released. + // #nosec G102 -- this test must reserve every interface used by the server. + portReservation, err := net.Listen("tcp", ":0") + require.NoError(t, err) + addr, ok := portReservation.Addr().(*net.TCPAddr) + require.True(t, ok) + port := addr.Port + require.NoError(t, portReservation.Close()) + + server := NewServer(port, config.DefaultMasterConfig()) + ctx, cancel := context.WithCancel(t.Context()) + startErr := make(chan error, 1) + + go func() { + startErr <- server.Start(ctx) + }() + + select { + case <-server.Ready(): + case err := <-startErr: + t.Fatalf("Start returned before reporting ready: %v", err) + case <-time.After(time.Second): + t.Fatal("server did not report ready after listening") + } + + cancel() + + select { + case err = <-startErr: + require.NoError(t, err) + case <-time.After(2 * time.Second): + t.Fatal("Start did not return after cancellation") + } + + // #nosec G102 -- this test verifies that shutdown released every interface. + rebound, err := net.Listen("tcp", fmt.Sprintf(":%d", port)) + require.NoError(t, err, "Start returned before releasing the listening port") + require.NoError(t, rebound.Close()) +} diff --git a/internal/mod/http/handler.go b/internal/mod/http/handler.go index 070ca49..3db853a 100644 --- a/internal/mod/http/handler.go +++ b/internal/mod/http/handler.go @@ -49,8 +49,21 @@ func NewHttpHandler(reqClient api.RequestClient, confManager config.ConfManager, } func (c *CustomHttpHandler) Run(ctx context.Context) { - httpConfig := c.confManager.LoadWithRemote().HttpServer - serverPort := httpConfig.Port + appConfig, err := c.confManager.LoadWithRemote() + if err != nil { + if appConfig == nil { + log.Errorf("Unable to load HTTP server config: %v", err) + select { + case c.errChan <- err: + case <-ctx.Done(): + default: + log.Warnf("Error channel is unavailable, dropping error: %v", err) + } + return + } + log.Warnf("Unable to reload HTTP server config, using last-known-good config: %v", err) + } + serverPort := appConfig.HttpServer.Port router := mux.NewRouter() router.HandleFunc("/ruleEngine/messages", server.RulesHandler(c.pubSub, c.queuedBytes)).Methods("POST") diff --git a/internal/mod/http/server/config.go b/internal/mod/http/server/config.go index 84fad54..68c2d5f 100644 --- a/internal/mod/http/server/config.go +++ b/internal/mod/http/server/config.go @@ -67,7 +67,15 @@ func LogConfigHandler() func(w http.ResponseWriter, r *http.Request) { func CurrentConfigHandler(confManager config.ConfManager) func(w http.ResponseWriter, r *http.Request) { return func(w http.ResponseWriter, r *http.Request) { - appConfig := confManager.LoadWithRemote() + appConfig, loadErr := confManager.LoadWithRemote() + if loadErr != nil { + if appConfig == nil { + log.Errorf("Failed to load current config: %v", loadErr) + http.Error(w, loadErr.Error(), http.StatusInternalServerError) + return + } + log.Warnf("Failed to reload current config, serving last-known-good config: %v", loadErr) + } bytes, err := json.Marshal(appConfig) if err != nil { diff --git a/internal/mod/rule/engine.go b/internal/mod/rule/engine.go index b9b5275..ed9eabc 100644 --- a/internal/mod/rule/engine.go +++ b/internal/mod/rule/engine.go @@ -16,6 +16,7 @@ package rule import ( "encoding/json" + stderrors "errors" "strconv" "sync" "time" @@ -51,6 +52,8 @@ type Engine struct { activeTopics mapset.Set[string] ruleDebounceTime map[string]*time.Time + + publishCollectInfoFn func(string) error } // UpdateRules updates the rules in the rule engine. @@ -169,9 +172,10 @@ func (e *Engine) getDeviceName() string { return e.deviceName } -// ConsumeNext shows how to process a message through the rule engine. -func (e *Engine) ConsumeNext(item rule_engine.RuleItem) { +// ConsumeNext processes a message through the rule engine. +func (e *Engine) ConsumeNext(item rule_engine.RuleItem) error { log.Debugf("consuming message: %+v", item) + var consumeErr error e.mu.RLock() rules := e.rules @@ -225,10 +229,22 @@ func (e *Engine) ConsumeNext(item rule_engine.RuleItem) { for _, action := range rule.Actions { if err := action.Run(curActivation, additionalArgs); err != nil { ruleLog.WithField("action", action.Name).Errorf("failed to run action: %v", err) + consumeErr = stderrors.Join( + consumeErr, + errors.Wrapf(err, "run action %s", action.Name), + ) } } - if err := model.PublishCollectInfo(collectInfoId); err != nil { + publishCollectInfo := e.publishCollectInfoFn + if publishCollectInfo == nil { + publishCollectInfo = model.PublishCollectInfo + } + if err := publishCollectInfo(collectInfoId); err != nil { ruleLog.Errorf("failed to publish collect info: %v", err) + consumeErr = stderrors.Join( + consumeErr, + errors.Wrap(err, "publish collect info"), + ) } else { ruleLog.Infof("published collect info") } @@ -236,6 +252,7 @@ func (e *Engine) ConsumeNext(item rule_engine.RuleItem) { } e.cleanupDebounceTime() + return consumeErr } func debounceKeyForRule(rule *rule_engine.Rule) string { diff --git a/internal/mod/rule/engine_test.go b/internal/mod/rule/engine_test.go index 25c108c..6b6026f 100644 --- a/internal/mod/rule/engine_test.go +++ b/internal/mod/rule/engine_test.go @@ -15,6 +15,7 @@ package rule import ( + "errors" "slices" "testing" "time" @@ -24,6 +25,11 @@ import ( mapset "github.com/deckarep/golang-set/v2" ) +var ( + errTestActionFailed = errors.New("action failed") + errTestPublishFailed = errors.New("publish failed") +) + func TestUpdateRulesKeepsAllTopicsActiveWhenAnyRuleHasNoTopic(t *testing.T) { t.Parallel() @@ -105,6 +111,9 @@ func TestConsumeNextDebouncesEachScopeSeparately(t *testing.T) { } engine := Engine{ ruleDebounceTime: make(map[string]*time.Time), + publishCollectInfoFn: func(string) error { + return nil + }, rules: []*rule_engine.Rule{ rule_engine.NewRule( []rule_engine.Condition{*condition}, @@ -125,8 +134,20 @@ func TestConsumeNextDebouncesEachScopeSeparately(t *testing.T) { }, } - engine.ConsumeNext(rule_engine.RuleItem{Msg: map[string]interface{}{"code": "A"}, Topic: "/fault", Ts: 1}) - engine.ConsumeNext(rule_engine.RuleItem{Msg: map[string]interface{}{"code": "B"}, Topic: "/fault", Ts: 2}) + if err := engine.ConsumeNext(rule_engine.RuleItem{ + Msg: map[string]interface{}{"code": "A"}, + Topic: "/fault", + Ts: 1, + }); err != nil { + t.Fatalf("consume scope A: %v", err) + } + if err := engine.ConsumeNext(rule_engine.RuleItem{ + Msg: map[string]interface{}{"code": "B"}, + Topic: "/fault", + Ts: 2, + }); err != nil { + t.Fatalf("consume scope B: %v", err) + } if len(triggered) != 2 { t.Fatalf("triggered scopes = %v, want both scopes to trigger independently", triggered) @@ -136,6 +157,93 @@ func TestConsumeNextDebouncesEachScopeSeparately(t *testing.T) { } } +func TestConsumeNextReturnsActionError(t *testing.T) { + t.Parallel() + + action, err := rule_engine.NewAction( + "failing-action", + map[string]interface{}{}, + func(map[string]interface{}) error { + return errTestActionFailed + }, + ) + if err != nil { + t.Fatalf("new action: %v", err) + } + publishCalls := 0 + engine := Engine{ + rules: []*rule_engine.Rule{testRuntimeRule(t, action)}, + ruleDebounceTime: make(map[string]*time.Time), + publishCollectInfoFn: func(string) error { + publishCalls++ + return nil + }, + } + + gotErr := engine.ConsumeNext(rule_engine.RuleItem{ + Msg: map[string]interface{}{}, + Topic: "/fault", + Ts: 1, + }) + + if !errors.Is(gotErr, errTestActionFailed) { + t.Fatalf("ConsumeNext error = %v, want action error", gotErr) + } + if publishCalls != 1 { + t.Fatalf("publish calls = %d, want 1", publishCalls) + } +} + +func TestConsumeNextReturnsPublishError(t *testing.T) { + t.Parallel() + + action, err := rule_engine.NewAction( + "successful-action", + map[string]interface{}{}, + rule_engine.EmptyActionImpl, + ) + if err != nil { + t.Fatalf("new action: %v", err) + } + engine := Engine{ + rules: []*rule_engine.Rule{testRuntimeRule(t, action)}, + ruleDebounceTime: make(map[string]*time.Time), + publishCollectInfoFn: func(string) error { + return errTestPublishFailed + }, + } + + gotErr := engine.ConsumeNext(rule_engine.RuleItem{ + Msg: map[string]interface{}{}, + Topic: "/fault", + Ts: 1, + }) + + if !errors.Is(gotErr, errTestPublishFailed) { + t.Fatalf("ConsumeNext error = %v, want publish error", gotErr) + } +} + +func testRuntimeRule(t *testing.T, action *rule_engine.Action) *rule_engine.Rule { + t.Helper() + + condition, err := rule_engine.NewCondition("true") + if err != nil { + t.Fatalf("new condition: %v", err) + } + return rule_engine.NewRule( + []rule_engine.Condition{*condition}, + []rule_engine.Action{*action}, + nil, + mapset.NewSet[string]("/fault"), + 0, + map[string]interface{}{ + "rule_name": "runtime-rule", + "rule_display_name": "runtime rule", + }, + ) +} + func testDiagnosisRule(id string, activeTopic string) *resources.DiagnosisRule { return &resources.DiagnosisRule{ Name: "projects/project-1/diagnosisRuleSets/set-1/diagnosisRules/" + id, diff --git a/internal/mod/rule/file_handlers/cancellation_test.go b/internal/mod/rule/file_handlers/cancellation_test.go new file mode 100644 index 0000000..7e06fa7 --- /dev/null +++ b/internal/mod/rule/file_handlers/cancellation_test.go @@ -0,0 +1,154 @@ +// Copyright 2026 coScene +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package file_handlers + +import ( + "context" + "os" + "path/filepath" + "testing" + "time" + + "github.com/coscene-io/coscout/pkg/rule_engine" + mapset "github.com/deckarep/golang-set/v2" + "github.com/stretchr/testify/require" +) + +func TestLogHandlerCancellationUnblocksFullRuleItemChannel(t *testing.T) { + t.Parallel() + + logPath := filepath.Join(t.TempDir(), "application.log") + require.NoError(t, os.WriteFile( + logPath, + []byte("2025-01-01 00:00:00.000 INFO started\n"), + 0o600, + )) + + ruleItems := make(chan rule_engine.RuleItem) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + done := make(chan struct{}) + result := make(chan error, 1) + go func() { + defer close(done) + result <- NewLogHandler().SendRuleItems( + ctx, + logPath, + mapset.NewSet[string](), + sendRuleItemToChannel(ctx, ruleItems), + ) + }() + + select { + case <-done: + t.Fatal("log handler returned before its blocked channel send was cancelled") + case <-time.After(50 * time.Millisecond): + } + + cancel() + select { + case err := <-result: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("log handler did not return after cancellation") + } + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) +} + +func TestHandlersReturnErrorsForUnreadableFiles(t *testing.T) { + t.Parallel() + + missingPath := filepath.Join(t.TempDir(), "missing") + ruleItems := make(chan rule_engine.RuleItem, 1) + activeTopics := mapset.NewSet[string]() + + for name, handler := range map[string]Interface{ + "log": NewLogHandler(), + "mcap": NewMcapHandler(), + "ros1": NewRos1Handler(), + } { + t.Run(name, func(t *testing.T) { + t.Parallel() + + err := handler.SendRuleItems( + t.Context(), + missingPath, + activeTopics, + sendRuleItemToChannel(t.Context(), ruleItems), + ) + require.Error(t, err) + }) + } +} + +func TestIntentionalSkipsReturnSuccess(t *testing.T) { + t.Parallel() + + ruleItems := make(chan rule_engine.RuleItem, 1) + + require.NoError(t, NewDefaultHandler().SendRuleItems( + t.Context(), + "unsupported.file", + mapset.NewSet[string](), + sendRuleItemToChannel(t.Context(), ruleItems), + )) + require.NoError(t, NewLogHandler().SendRuleItems( + t.Context(), + "missing.log", + mapset.NewSet[string]("/some_other_topic"), + sendRuleItemToChannel(t.Context(), ruleItems), + )) +} + +func TestLogHandlerEOFReturnsSuccess(t *testing.T) { + t.Parallel() + + logPath := filepath.Join(t.TempDir(), "application.log") + require.NoError(t, os.WriteFile( + logPath, + []byte("2025-01-01 00:00:00.000 INFO complete\n"), + 0o600, + )) + ruleItems := make(chan rule_engine.RuleItem, 1) + + require.NoError(t, NewLogHandler().SendRuleItems( + t.Context(), + logPath, + mapset.NewSet[string](), + sendRuleItemToChannel(t.Context(), ruleItems), + )) + require.Len(t, ruleItems, 1) +} + +func sendRuleItemToChannel( + ctx context.Context, + ruleItems chan<- rule_engine.RuleItem, +) func(rule_engine.RuleItem) bool { + return func(item rule_engine.RuleItem) bool { + select { + case ruleItems <- item: + return true + case <-ctx.Done(): + return false + } + } +} diff --git a/internal/mod/rule/file_handlers/default.go b/internal/mod/rule/file_handlers/default.go index 13da61b..1ef4798 100644 --- a/internal/mod/rule/file_handlers/default.go +++ b/internal/mod/rule/file_handlers/default.go @@ -15,6 +15,7 @@ package file_handlers import ( + "context" "os" "strings" "time" @@ -69,8 +70,14 @@ func (h *defaultHandler) GetStartTimeEndTime(filePath string) (*time.Time, *time return &birthTime, &modTime, nil } -func (h *defaultHandler) SendRuleItems(filePath string, activeTopics mapset.Set[string], sendRuleItem func(rule_engine.RuleItem) bool) { +func (h *defaultHandler) SendRuleItems( + _ context.Context, + filePath string, + _ mapset.Set[string], + _ func(rule_engine.RuleItem) bool, +) error { log.Info("filePath not supported by rule engine: ", filePath, ", skip sending rule items") + return nil } func (h *defaultHandler) IsFinished(filePath string) bool { diff --git a/internal/mod/rule/file_handlers/interface.go b/internal/mod/rule/file_handlers/interface.go index c61e2d8..16680ae 100644 --- a/internal/mod/rule/file_handlers/interface.go +++ b/internal/mod/rule/file_handlers/interface.go @@ -15,6 +15,7 @@ package file_handlers import ( + "context" "os" "time" @@ -36,8 +37,14 @@ type Interface interface { // IsFinished checks if the file is completely written and no more updates are expected. IsFinished(filePath string) bool - // SendRuleItems sends rule items until the file ends or sendRuleItem returns false. - SendRuleItems(filePath string, activeTopics mapset.Set[string], sendRuleItem func(rule_engine.RuleItem) bool) + // SendRuleItems sends rule items until the file ends, the context is + // cancelled, or sendRuleItem returns false. + SendRuleItems( + ctx context.Context, + filePath string, + activeTopics mapset.Set[string], + sendRuleItem func(rule_engine.RuleItem) bool, + ) error } // defaultGetFileSize provides default implementations for some methods. diff --git a/internal/mod/rule/file_handlers/log.go b/internal/mod/rule/file_handlers/log.go index b48afd8..f42c2c7 100644 --- a/internal/mod/rule/file_handlers/log.go +++ b/internal/mod/rule/file_handlers/log.go @@ -15,6 +15,7 @@ package file_handlers import ( + "context" "os" "strings" "time" @@ -75,35 +76,44 @@ func (h *logHandler) GetStartTimeEndTime(filePath string) (*time.Time, *time.Tim return startTime, endTime, nil } -func (h *logHandler) SendRuleItems(filePath string, activeTopics mapset.Set[string], sendRuleItem func(rule_engine.RuleItem) bool) { +func (h *logHandler) SendRuleItems( + ctx context.Context, + filePath string, + activeTopics mapset.Set[string], + sendRuleItem func(rule_engine.RuleItem) bool, +) error { if activeTopics.Cardinality() > 0 && !activeTopics.Contains("/external_log") { - return + return nil } logFileReader, err := os.Open(filePath) if err != nil { log.Errorf("open log file [%s] failed: %v", filePath, err) - return + return errors.Errorf("open log file [%s]: %v", filePath, err) } defer logFileReader.Close() reader, err := log_reader.NewLogReader(logFileReader, filePath) if err != nil { log.Errorf("failed to create log reader for log file %s: %v", filePath, err) - return + return errors.Errorf("create log reader for log file %s: %v", filePath, err) } iter, err := log_reader.NewLogIterator(reader) if err != nil { log.Errorf("failed to create log iterator for log file %s: %v", filePath, err) - return + return errors.Errorf("create log iterator for log file %s: %v", filePath, err) } for { + if ctx.Err() != nil { + return ctx.Err() + } + stampedLog, hasNext := iter.Next() if !hasNext { log.Infof("finished sending rule items for log file %s", filePath) - break + return nil } sec := stampedLog.Timestamp.Unix() @@ -124,7 +134,10 @@ func (h *logHandler) SendRuleItems(filePath string, activeTopics mapset.Set[stri Ts: tsFloat, Source: filePath, }) { - return + if ctx.Err() != nil { + return ctx.Err() + } + return nil } } } diff --git a/internal/mod/rule/file_handlers/mcap.go b/internal/mod/rule/file_handlers/mcap.go index 5d6fa19..02d67d5 100644 --- a/internal/mod/rule/file_handlers/mcap.go +++ b/internal/mod/rule/file_handlers/mcap.go @@ -16,6 +16,7 @@ package file_handlers import ( "bytes" + "context" "encoding/json" "io" "os" @@ -166,25 +167,30 @@ func getStartTimeEndTimeForUncompletedMcap(filePath string) (*time.Time, *time.T return &start, &end, nil } -func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[string], sendRuleItem func(rule_engine.RuleItem) bool) { +func (h *mcapHandler) SendRuleItems( + ctx context.Context, + filepath string, + activeTopics mapset.Set[string], + sendRuleItem func(rule_engine.RuleItem) bool, +) error { file, err := os.Open(filepath) if err != nil { log.Errorf("failed to open MCAP file [%s]: %v", filepath, err) - return + return errors.Errorf("open MCAP file [%s]: %v", filepath, err) } defer file.Close() reader, err := mcap.NewReader(file) if err != nil { log.Errorf("failed to create MCAP reader: %v", err) - return + return errors.Errorf("create MCAP reader: %v", err) } defer reader.Close() info, err := reader.Info() if err != nil { log.Errorf("failed to get info for MCAP file [%s]: %v", filepath, err) - return + return errors.Errorf("get info for MCAP file [%s]: %v", filepath, err) } // targetTopics will be empty if no active topics are provided @@ -199,7 +205,7 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str if targetTopics.Cardinality() == 0 { log.Infof("no active topics found in MCAP file %s", filepath) - return + return nil } } log.Infof("sending rule items for MCAP file %s with topics: %v", filepath, targetTopics) @@ -207,7 +213,7 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str iter, err := reader.Messages(mcap.WithTopics(targetTopics.ToSlice())) if err != nil { log.Errorf("failed to create message iterator: %v", err) - return + return errors.Errorf("create MCAP message iterator: %v", err) } // Remove msgBuf since we'll decode directly to map where possible @@ -217,16 +223,24 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str ros2Decoders := make(map[string]mcap_ros2.DecoderFunction) descriptors := make(map[uint16]protoreflect.MessageDescriptor) errTopics := map[string]struct{}{} + logMessageError := func(err error) { + log.Error(err) + } for { + if ctx.Err() != nil { + return ctx.Err() + } + schema, channel, message, err := iter.NextInto(msg) if err != nil { if errors.Is(err, io.EOF) { log.Infof("finished sending rule items for MCAP file %s", filepath) + return nil } else { log.Errorf("error reading message: %v", err) + return errors.Errorf("read MCAP message: %v", err) } - return } if _, ok := errTopics[channel.Topic]; ok { @@ -241,11 +255,14 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str switch channel.MessageEncoding { case "json": if err := json.Unmarshal(message.Data, &decoded); err != nil { - log.Errorf("failed to unmarshal JSON: %v", err) + logMessageError(errors.Errorf("unmarshal JSON: %v", err)) continue } default: - log.Errorf("unsupported message encoding for schema-less channel: %s", channel.MessageEncoding) + logMessageError(errors.Errorf( + "unsupported message encoding for schema-less channel: %s", + channel.MessageEncoding, + )) continue } } else { @@ -256,7 +273,7 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str packageName := strings.Split(schema.Name, "/")[0] transcoder, err = ros1msg.NewJSONTranscoder(packageName, schema.Data) if err != nil { - log.Errorf("failed to create JSON transcoder: %v", err) + logMessageError(errors.Errorf("create JSON transcoder: %v", err)) continue } transcoders[channel.SchemaID] = transcoder @@ -265,11 +282,11 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str buf := &bytes.Buffer{} msgReader.Reset(message.Data) if err := transcoder.Transcode(buf, msgReader); err != nil { - log.Errorf("failed to transcode: %v", err) + logMessageError(errors.Errorf("transcode ros1 message: %v", err)) continue } if err := json.Unmarshal(buf.Bytes(), &decoded); err != nil { - log.Errorf("failed to unmarshal transcoded data: %v", err) + logMessageError(errors.Errorf("unmarshal transcoded data: %v", err)) continue } @@ -278,7 +295,7 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str if !ok { dynamicDecoders, err := mcap_ros2.GenerateDynamic(schema.Name, string(schema.Data)) if err != nil { - log.Errorf("failed to generate dynamic schema decoder: %v", err) + logMessageError(errors.Errorf("generate dynamic schema decoder: %v", err)) continue } @@ -289,7 +306,7 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str var decoderOk bool decoder, decoderOk = dynamicDecoders[schema.Name] if !decoderOk { - log.Errorf("failed to find decoder for schema: %s", schema.Name) + logMessageError(errors.Errorf("find decoder for schema: %s", schema.Name)) continue } } @@ -305,7 +322,11 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str decoded, err = decoder(message.Data) }() if panicErr != nil { - log.Errorf("failed to decode message: %v", panicErr) + logMessageError(errors.Errorf("decode message: %v", panicErr)) + continue + } + if err != nil { + logMessageError(errors.Errorf("decode message: %v", err)) continue } @@ -314,45 +335,53 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str if !ok { fileDescriptorSet := &descriptorpb.FileDescriptorSet{} if err := proto.Unmarshal(schema.Data, fileDescriptorSet); err != nil { - log.Errorf("failed to build file descriptor set: %v", err) + logMessageError(errors.Errorf("build file descriptor set: %v", err)) continue } files, err := protodesc.FileOptions{}.NewFiles(fileDescriptorSet) if err != nil { - log.Errorf("failed to create file descriptor: %v", err) + logMessageError(errors.Errorf("create file descriptor: %v", err)) continue } descriptor, err := files.FindDescriptorByName(protoreflect.FullName(schema.Name)) if err != nil { - log.Errorf("failed to find descriptor: %v", err) + logMessageError(errors.Errorf("find descriptor: %v", err)) + continue + } + var descriptorOk bool + messageDescriptor, descriptorOk = descriptor.(protoreflect.MessageDescriptor) + if !descriptorOk { + logMessageError(errors.Errorf( + "descriptor %s is not a message descriptor", + schema.Name, + )) continue } - messageDescriptor, _ = descriptor.(protoreflect.MessageDescriptor) descriptors[channel.SchemaID] = messageDescriptor } protoMsg := dynamicpb.NewMessage(messageDescriptor) if err := proto.Unmarshal(message.Data, protoMsg); err != nil { - log.Errorf("failed to parse protobuf message: %v", err) + logMessageError(errors.Errorf("parse protobuf message: %v", err)) continue } marshalledBytes, err := protojson.Marshal(protoMsg) if err != nil { - log.Errorf("failed to marshal protobuf to JSON: %v", err) + logMessageError(errors.Errorf("marshal protobuf to JSON: %v", err)) continue } if err := json.Unmarshal(marshalledBytes, &decoded); err != nil { - log.Errorf("failed to unmarshal protobuf JSON: %v", err) + logMessageError(errors.Errorf("unmarshal protobuf JSON: %v", err)) continue } case "jsonschema": if err := json.Unmarshal(message.Data, &decoded); err != nil { - log.Errorf("failed to unmarshal JSON: %v", err) + logMessageError(errors.Errorf("unmarshal JSON schema message: %v", err)) continue } default: - log.Errorf("unsupported schema encoding: %s", schema.Encoding) + logMessageError(errors.Errorf("unsupported schema encoding: %s", schema.Encoding)) continue } } @@ -363,7 +392,10 @@ func (h *mcapHandler) SendRuleItems(filepath string, activeTopics mapset.Set[str Topic: channel.Topic, Source: filepath, }) { - return + if ctx.Err() != nil { + return ctx.Err() + } + return nil } } } diff --git a/internal/mod/rule/file_handlers/ros1.go b/internal/mod/rule/file_handlers/ros1.go index 257e98c..d1bca36 100644 --- a/internal/mod/rule/file_handlers/ros1.go +++ b/internal/mod/rule/file_handlers/ros1.go @@ -16,6 +16,7 @@ package file_handlers import ( "bytes" + "context" "encoding/json" "os" "strings" @@ -82,24 +83,30 @@ func (h *ros1Handler) GetStartTimeEndTime(filePath string) (*time.Time, *time.Ti return &start, &end, nil } -func (h *ros1Handler) SendRuleItems(filePath string, activeTopics mapset.Set[string], sendRuleItem func(rule_engine.RuleItem) bool) { +func (h *ros1Handler) SendRuleItems( + ctx context.Context, + filePath string, + activeTopics mapset.Set[string], + sendRuleItem func(rule_engine.RuleItem) bool, +) error { reader, err := os.Open(filePath) if err != nil { log.Errorf("failed to open ros1 file [%s]: %v", filePath, err) - return + return errors.Errorf("open ros1 file [%s]: %v", filePath, err) } + defer reader.Close() // Create bag reader bagReader, err := rosbag.NewReader(reader) if err != nil { log.Errorf("failed to create bag reader: %v", err) - return + return errors.Errorf("create ros1 bag reader: %v", err) } info, err := bagReader.Info() if err != nil { log.Errorf("failed to get info for ros1 file %s: %v", filePath, err) - return + return errors.Errorf("get info for ros1 file %s: %v", filePath, err) } // targetTopics will be empty if there is no active topic @@ -115,7 +122,7 @@ func (h *ros1Handler) SendRuleItems(filePath string, activeTopics mapset.Set[str if targetTopics.Cardinality() == 0 { log.Infof("no active topics matched in ros1 file %s, skipping", filePath) - return + return nil } } log.Infof("sending rule items for ros1 file %s with topics: %v", filePath, targetTopics) @@ -124,7 +131,7 @@ func (h *ros1Handler) SendRuleItems(filePath string, activeTopics mapset.Set[str it, err := bagReader.Messages() if err != nil { log.Errorf("failed to create message iterator: %v", err) - return + return errors.Errorf("create ros1 message iterator: %v", err) } // Map to store JSON transcoders for each connection @@ -132,13 +139,20 @@ func (h *ros1Handler) SendRuleItems(filePath string, activeTopics mapset.Set[str // Set to store conn failed to transcode failedConnToTranscode := mapset.NewSet[uint32]() + logMessageError := func(err error) { + log.Error(err) + } // Read messages for it.More() { + if ctx.Err() != nil { + return ctx.Err() + } + conn, msg, err := it.Next() if err != nil { log.Errorf("failed to read message: %v", err) - continue + return errors.Errorf("read ros1 message: %v", err) } if targetTopics.Cardinality() > 0 && !targetTopics.Contains(conn.Topic) { continue @@ -158,7 +172,7 @@ func (h *ros1Handler) SendRuleItems(filePath string, activeTopics mapset.Set[str } transcoder, err = ros1msg.NewJSONTranscoder(parentPackage, conn.Data.MessageDefinition) if err != nil { - log.Errorf("failed to create JSON transcoder: %v", err) + logMessageError(errors.Errorf("create JSON transcoder: %v", err)) failedConnToTranscode.Add(conn.Conn) continue } @@ -169,14 +183,14 @@ func (h *ros1Handler) SendRuleItems(filePath string, activeTopics mapset.Set[str var buf bytes.Buffer err = transcoder.Transcode(&buf, bytes.NewReader(msg.Data)) if err != nil { - log.Errorf("failed to transcode message: %v", err) + logMessageError(errors.Errorf("transcode ros1 message: %v", err)) continue } // Parse JSON into structured data var structuredData map[string]interface{} if err := json.Unmarshal(buf.Bytes(), &structuredData); err != nil { - log.Errorf("failed to parse JSON: %v", err) + logMessageError(errors.Errorf("parse ros1 message JSON: %v", err)) continue } @@ -188,8 +202,12 @@ func (h *ros1Handler) SendRuleItems(filePath string, activeTopics mapset.Set[str Topic: conn.Topic, Source: filePath, }) { - return + if ctx.Err() != nil { + return ctx.Err() + } + return nil } } log.Infof("finished sending rule items for ros1 file %s", filePath) + return nil } diff --git a/internal/mod/rule/file_state_handler/file_state_handler.go b/internal/mod/rule/file_state_handler/file_state_handler.go index ea3efa4..7a4b8b9 100644 --- a/internal/mod/rule/file_state_handler/file_state_handler.go +++ b/internal/mod/rule/file_state_handler/file_state_handler.go @@ -31,7 +31,6 @@ import ( mapset "github.com/deckarep/golang-set/v2" "github.com/djherbis/times" "github.com/pkg/errors" - "github.com/samber/lo" log "github.com/sirupsen/logrus" ) @@ -54,6 +53,10 @@ type FileStateHandler interface { // to mark them as processed. MarkProcessedFile(filename string) error + // MarkFailedFile marks a file as permanently failed. Failed files remain + // terminal until their size or modification time changes. + MarkFailedFile(filename string) error + // GetFileHandler returns the handler for a given file path. GetFileHandler(filePath string) file_handlers.Interface } @@ -66,6 +69,7 @@ const ( processStateSeenOnce processStateReadyToProcess processStateProcessed + processStateFailed ) // FileState represents the state of a file in the system. @@ -94,7 +98,8 @@ type SavedState struct { // fileStateHandler is used to keep track of the state of files in directories. type fileStateHandler struct { state map[string]FileState - updateLock sync.Mutex + stateLock sync.RWMutex + saveLock sync.Mutex listenDirs mapset.Set[string] collectDirs mapset.Set[string] activeTopics mapset.Set[string] @@ -187,27 +192,77 @@ func (f *fileStateHandler) loadState() error { log.Warnf("Invalid file state file: %v, reset to init", err) } + if savedState.State == nil { + savedState.State = make(map[string]FileState) + } + listenDirs := mapset.NewSet(savedState.ListenDirs...) + collectDirs := mapset.NewSet(savedState.CollectDirs...) + + f.stateLock.Lock() f.state = savedState.State - f.listenDirs = mapset.NewSet[string]() - f.collectDirs = mapset.NewSet[string]() + f.listenDirs = listenDirs + f.collectDirs = collectDirs + f.stateLock.Unlock() + + return nil +} + +func (f *fileStateHandler) savedStateSnapshot() SavedState { + f.stateLock.RLock() + defer f.stateLock.RUnlock() - for _, dir := range savedState.ListenDirs { - f.listenDirs.Add(dir) + state := make(map[string]FileState, len(f.state)) + for filename, fileState := range f.state { + state[filename] = fileState + } + + return SavedState{ + State: state, + ListenDirs: f.listenDirs.ToSlice(), + CollectDirs: f.collectDirs.ToSlice(), } - for _, dir := range savedState.CollectDirs { - f.collectDirs.Add(dir) +} + +func (f *fileStateHandler) stateKeysSnapshot() []string { + f.stateLock.RLock() + defer f.stateLock.RUnlock() + + filenames := make([]string, 0, len(f.state)) + for filename := range f.state { + filenames = append(filenames, filename) } + return filenames +} - return nil +func (f *fileStateHandler) stateSnapshot() map[string]FileState { + f.stateLock.RLock() + defer f.stateLock.RUnlock() + + state := make(map[string]FileState, len(f.state)) + for filename, fileState := range f.state { + state[filename] = fileState + } + return state +} + +func (f *fileStateHandler) listenDirsSnapshot() mapset.Set[string] { + f.stateLock.RLock() + defer f.stateLock.RUnlock() + return mapset.NewSet(f.listenDirs.ToSlice()...) +} + +func (f *fileStateHandler) collectDirsSnapshot() mapset.Set[string] { + f.stateLock.RLock() + defer f.stateLock.RUnlock() + return mapset.NewSet(f.collectDirs.ToSlice()...) } // SaveState saves the current state to disk. func (f *fileStateHandler) saveState() error { - savedState := SavedState{ - State: f.state, - ListenDirs: f.listenDirs.ToSlice(), - CollectDirs: f.collectDirs.ToSlice(), - } + f.saveLock.Lock() + defer f.saveLock.Unlock() + + savedState := f.savedStateSnapshot() data, err := json.MarshalIndent(savedState, "", " ") if err != nil { @@ -227,15 +282,34 @@ func (f *fileStateHandler) saveState() error { // setFileState updates the state for a given file path. func (f *fileStateHandler) setFileState(filePath string, state FileState) { + if err := f.updateFileState(filePath, func(FileState, bool) (FileState, bool) { + return state, true + }); err != nil { + log.Errorf("Failed to update file state for %s: %v", filePath, err) + } +} + +// updateFileState applies a read-modify-write operation while holding stateLock. +// The callback must only update in-memory state and must not perform file I/O. +func (f *fileStateHandler) updateFileState( + filePath string, + update func(current FileState, exists bool) (next FileState, keep bool), +) error { absPath, err := filepath.Abs(filePath) if err != nil { - log.Errorf("Failed to get absolute path for %s: %v", filePath, err) - return + return errors.Errorf("failed to get absolute path for %s: %v", filePath, err) } - f.updateLock.Lock() - defer f.updateLock.Unlock() - f.state[absPath] = state + f.stateLock.Lock() + defer f.stateLock.Unlock() + current, exists := f.state[absPath] + next, keep := update(current, exists) + if !keep { + delete(f.state, absPath) + return nil + } + f.state[absPath] = next + return nil } // delFileState removes the state for a given file path. @@ -245,8 +319,8 @@ func (f *fileStateHandler) delFileState(filePath string) error { return errors.Errorf("failed to get absolute path for %s: %v", filePath, err) } - f.updateLock.Lock() - defer f.updateLock.Unlock() + f.stateLock.Lock() + defer f.stateLock.Unlock() delete(f.state, absPath) return nil } @@ -259,6 +333,8 @@ func (f *fileStateHandler) getFileState(filePath string) (FileState, bool) { return FileState{}, false } + f.stateLock.RLock() + defer f.stateLock.RUnlock() state, exists := f.state[absPath] return state, exists } @@ -278,11 +354,24 @@ func (f *fileStateHandler) checkpointFileState(checkpoints map[string]fileStateC func (f *fileStateHandler) restoreFileStates(checkpoints map[string]fileStateCheckpoint) { for filePath, checkpoint := range checkpoints { - if checkpoint.exists { - f.setFileState(filePath, checkpoint.state) - continue - } - if err := f.delFileState(filePath); err != nil { + err := f.updateFileState(filePath, func(current FileState, exists bool) (FileState, bool) { + if checkpoint.exists { + restored := checkpoint.state + if exists && current.ProcessState != checkpoint.state.ProcessState { + restored.ProcessState = current.ProcessState + } + return restored, true + } + + // A non-default process state can only have been written after the + // checkpoint. Preserve that concurrent transition instead of + // deleting the newly-created state. + if exists && current.ProcessState != processStateUnprocessed { + return current, true + } + return FileState{}, false + }) + if err != nil { log.Errorf("Failed to roll back file state for %s: %v", filePath, err) } } @@ -290,7 +379,7 @@ func (f *fileStateHandler) restoreFileStates(checkpoints map[string]fileStateChe // updateDeletedFileState removes state entries for files that no longer exist. func (f *fileStateHandler) updateDeletedFileState() error { - for _, filename := range lo.Keys(f.state) { + for _, filename := range f.stateKeysSnapshot() { if _, err := os.Stat(filename); os.IsNotExist(err) { if err = f.delFileState(filename); err != nil { return errors.Errorf("failed to delete file state for %s: %v", filename, err) @@ -313,10 +402,10 @@ func (f *fileStateHandler) UpdateListenDirs(conf config.DefaultModConfConfig) er log.Infof("Start updating directories, listen_dirs: %v", newListenDirs) // Find directories that are no longer being monitored - deleteDirs := f.listenDirs.Difference(newListenDirs) + deleteDirs := f.listenDirsSnapshot().Difference(newListenDirs) // Delete state of files in directories that are no longer scanned - for filename := range f.state { + for _, filename := range f.stateKeysSnapshot() { for dir := range deleteDirs.Iter() { if strings.HasPrefix(filename, dir) { if err := f.delFileState(filename); err != nil { @@ -400,7 +489,9 @@ func (f *fileStateHandler) UpdateListenDirs(conf config.DefaultModConfConfig) er } // Update directory sets + f.stateLock.Lock() f.listenDirs = newListenDirs + f.stateLock.Unlock() // Clean up deleted files if err := f.updateDeletedFileState(); err != nil { @@ -427,10 +518,10 @@ func (f *fileStateHandler) UpdateCollectDirs(whitelist []string, conf config.Def log.Infof("Start updating directories, collector_dirs: %v", newCollectDirs) // Find directories that are no longer being monitored - deleteDirs := f.collectDirs.Difference(newCollectDirs) + deleteDirs := f.collectDirsSnapshot().Difference(newCollectDirs) // Delete state of files in directories that are no longer collected - for filename := range f.state { + for _, filename := range f.stateKeysSnapshot() { for dir := range deleteDirs.Iter() { if strings.HasPrefix(filename, dir) { if err := f.delFileState(filename); err != nil { @@ -547,7 +638,9 @@ func (f *fileStateHandler) UpdateCollectDirs(whitelist []string, conf config.Def } // Update directory sets + f.stateLock.Lock() f.collectDirs = newCollectDirs + f.stateLock.Unlock() // Clean up deleted files if err := f.updateDeletedFileState(); err != nil { @@ -563,6 +656,39 @@ func (f *fileStateHandler) UpdateCollectDirs(whitelist []string, conf config.Def return nil } +func fileContentsChanged(state FileState, info os.FileInfo) bool { + return state.Size != info.Size() || state.ModifyTime != info.ModTime().UnixNano() +} + +func nextScannedProcessState(state FileState, exists bool, info os.FileInfo) processState { + if exists && state.ProcessState == processStateFailed && fileContentsChanged(state, info) { + return processStateUnprocessed + } + if exists { + return state.ProcessState + } + return processStateUnprocessed +} + +// applyScannedFileState merges a scan result with the latest in-memory state. +// If processing advanced while the scanner was doing file I/O, the latest +// process state wins over the scanner's stale observation. +func (f *fileStateHandler) applyScannedFileState( + filePath string, + observed FileState, + observedExists bool, + scanned FileState, +) { + if err := f.updateFileState(filePath, func(current FileState, exists bool) (FileState, bool) { + if exists && (!observedExists || current.ProcessState != observed.ProcessState) { + scanned.ProcessState = current.ProcessState + } + return scanned, true + }); err != nil { + log.Errorf("Failed to apply scanned file state for %s: %v", filePath, err) + } +} + func (f *fileStateHandler) processListenFile(absPath string, info os.FileInfo, skipPeriodHours int) { fileState, hasFileState := f.getFileState(absPath) @@ -570,26 +696,34 @@ func (f *fileStateHandler) processListenFile(absPath string, info os.FileInfo, s if hasFileState && !fileState.Unsupported && fileState.ProcessState == processStateProcessed { return } + // A permanently failed file is retried only after its contents change. + if hasFileState && fileState.ProcessState == processStateFailed && !fileContentsChanged(fileState, info) { + return + } handler := f.GetFileHandler(absPath) if handler == nil { // No handler supported for file, mark as unsupported if not already if !hasFileState || !fileState.Unsupported { - f.setFileState(absPath, FileState{ - Size: info.Size(), - Unsupported: true, - }) + newState := fileState + newState.Size = info.Size() + newState.ModifyTime = info.ModTime().UnixNano() + newState.Unsupported = true + newState.ProcessState = nextScannedProcessState(fileState, hasFileState, info) + f.applyScannedFileState(absPath, fileState, hasFileState, newState) } return } // Check if file is too old if !info.IsDir() && info.ModTime().Before(time.Now().Add(-time.Duration(skipPeriodHours)*time.Hour)) { - f.setFileState(absPath, FileState{ - Size: info.Size(), - Unsupported: true, - TooOld: true, - }) + newState := fileState + newState.Size = info.Size() + newState.ModifyTime = info.ModTime().UnixNano() + newState.Unsupported = true + newState.TooOld = true + newState.ProcessState = nextScannedProcessState(fileState, hasFileState, info) + f.applyScannedFileState(absPath, fileState, hasFileState, newState) return } @@ -601,11 +735,12 @@ func (f *fileStateHandler) processListenFile(absPath string, info os.FileInfo, s } newState.Size = info.Size() - newState.ModifyTime = info.ModTime().Unix() + newState.ModifyTime = info.ModTime().UnixNano() newState.IsListening = true newState.Unsupported = false newState.TooOld = false - f.setFileState(absPath, newState) + newState.ProcessState = nextScannedProcessState(fileState, hasFileState, info) + f.applyScannedFileState(absPath, fileState, hasFileState, newState) return } else { log.Infof("Listening file %s is not finished, will be processed later", absPath) @@ -615,27 +750,28 @@ func (f *fileStateHandler) processListenFile(absPath string, info os.FileInfo, s func (f *fileStateHandler) processCollectFile(absPath string, info os.FileInfo) { fileState, hasFileState := f.getFileState(absPath) - // Skip file if file is not modified + // Skip terminal files while their contents are unchanged. if hasFileState && !fileState.Unsupported && - fileState.ProcessState == processStateProcessed && - fileState.ModifyTime == info.ModTime().Unix() { + (fileState.ProcessState == processStateProcessed || fileState.ProcessState == processStateFailed) && + !fileContentsChanged(fileState, info) { return } handler := f.GetFileHandler(absPath) if handler == nil { - f.setFileState(absPath, FileState{ - Size: info.Size(), - ModifyTime: info.ModTime().Unix(), - Unsupported: true, - }) + newState := fileState + newState.Size = info.Size() + newState.ModifyTime = info.ModTime().UnixNano() + newState.Unsupported = true + newState.ProcessState = nextScannedProcessState(fileState, hasFileState, info) + f.applyScannedFileState(absPath, fileState, hasFileState, newState) return } // Update state newState := fileState - if !hasFileState || fileState.ModifyTime != info.ModTime().Unix() { + if !hasFileState || fileContentsChanged(fileState, info) { var startTime, endTime *time.Time // Try to use file creation time and modification time if available @@ -654,11 +790,12 @@ func (f *fileStateHandler) processCollectFile(absPath string, info os.FileInfo) if err != nil { log.Errorf("Error getting start and end time for %s: %v, skip collect", absPath, err) - f.setFileState(absPath, FileState{ - Size: info.Size(), - ModifyTime: info.ModTime().Unix(), - Unsupported: true, - }) + newState := fileState + newState.Size = info.Size() + newState.ModifyTime = info.ModTime().UnixNano() + newState.Unsupported = true + newState.ProcessState = nextScannedProcessState(fileState, hasFileState, info) + f.applyScannedFileState(absPath, fileState, hasFileState, newState) return } } @@ -667,7 +804,7 @@ func (f *fileStateHandler) processCollectFile(absPath string, info os.FileInfo) Size: info.Size(), StartTime: startTime.Unix(), EndTime: endTime.Unix(), - ModifyTime: info.ModTime().Unix(), + ModifyTime: info.ModTime().UnixNano(), } log.Infof("Collecting file %s start time: %s, end time: %s", absPath, startTime.UTC().String(), endTime.UTC().String()) @@ -675,26 +812,31 @@ func (f *fileStateHandler) processCollectFile(absPath string, info os.FileInfo) newState.IsCollecting = true newState.ProcessState = processStateProcessed + if fileState.ProcessState == processStateFailed && fileContentsChanged(fileState, info) { + newState.ProcessState = processStateUnprocessed + } - f.setFileState(absPath, newState) + f.applyScannedFileState(absPath, fileState, hasFileState, newState) } func (f *fileStateHandler) UpdateFilesProcessState() error { - for filename, state := range f.state { - if !state.IsListening { - continue - } + for _, filename := range f.stateKeysSnapshot() { + if err := f.updateFileState(filename, func(state FileState, exists bool) (FileState, bool) { + if !exists || !state.IsListening { + return state, exists + } - switch state.ProcessState { - case processStateUnprocessed: - state.ProcessState = processStateSeenOnce - case processStateSeenOnce: - state.ProcessState = processStateReadyToProcess - case processStateReadyToProcess, processStateProcessed: - continue + switch state.ProcessState { + case processStateUnprocessed: + state.ProcessState = processStateSeenOnce + case processStateSeenOnce: + state.ProcessState = processStateReadyToProcess + case processStateReadyToProcess, processStateProcessed, processStateFailed: + } + return state, true + }); err != nil { + return errors.Errorf("failed to update process state for %s: %v", filename, err) } - - f.setFileState(filename, state) } if err := f.saveState(); err != nil { return errors.Errorf("failed to save state: %v", err) @@ -704,15 +846,29 @@ func (f *fileStateHandler) UpdateFilesProcessState() error { } func (f *fileStateHandler) MarkProcessedFile(filename string) error { - state, exists := f.getFileState(filename) + return f.markFileProcessState(filename, processStateProcessed) +} + +func (f *fileStateHandler) MarkFailedFile(filename string) error { + return f.markFileProcessState(filename, processStateFailed) +} + +func (f *fileStateHandler) markFileProcessState(filename string, nextProcessState processState) error { + exists := false + if err := f.updateFileState(filename, func(state FileState, currentExists bool) (FileState, bool) { + if !currentExists { + return state, false + } + exists = true + state.ProcessState = nextProcessState + return state, true + }); err != nil { + return err + } if !exists { return errors.Errorf("file state for %s does not exist", filename) } - state.ProcessState = processStateProcessed - - f.setFileState(filename, state) - if err := f.saveState(); err != nil { return errors.Errorf("failed to save state: %v", err) } @@ -725,11 +881,8 @@ type FileFilter func(string, FileState) bool // Files returns files matching the given filters. func (f *fileStateHandler) Files(filters ...FileFilter) []FileState { - f.updateLock.Lock() - defer f.updateLock.Unlock() - var result []FileState - for filename, state := range f.state { + for filename, state := range f.stateSnapshot() { if state.Unsupported { continue } diff --git a/internal/mod/rule/file_state_handler/file_state_handler_test.go b/internal/mod/rule/file_state_handler/file_state_handler_test.go index 0bc7124..cbda5a0 100644 --- a/internal/mod/rule/file_state_handler/file_state_handler_test.go +++ b/internal/mod/rule/file_state_handler/file_state_handler_test.go @@ -15,15 +15,80 @@ package file_state_handler import ( + "context" + "fmt" "os" "path/filepath" "reflect" + "sync" "testing" + "time" "github.com/coscene-io/coscout/internal/config" + "github.com/coscene-io/coscout/internal/mod/rule/file_handlers" + "github.com/coscene-io/coscout/pkg/rule_engine" mapset "github.com/deckarep/golang-set/v2" + "github.com/stretchr/testify/require" ) +type blockingFinishedHandler struct { + started chan struct{} + release chan struct{} +} + +type fileInfoWithModTime struct { + os.FileInfo + modTime time.Time +} + +func (i fileInfoWithModTime) ModTime() time.Time { + return i.modTime +} + +func (h *blockingFinishedHandler) CheckFilePath(string) bool { + return true +} + +func (h *blockingFinishedHandler) GetStartTimeEndTime(string) (*time.Time, *time.Time, error) { + now := time.Now() + return &now, &now, nil +} + +func (h *blockingFinishedHandler) GetFileSize(string) (int64, error) { + return 0, nil +} + +func (h *blockingFinishedHandler) IsFinished(string) bool { + if h.started != nil { + close(h.started) + } + if h.release != nil { + <-h.release + } + return true +} + +func (h *blockingFinishedHandler) SendRuleItems( + context.Context, + string, + mapset.Set[string], + func(rule_engine.RuleItem) bool, +) error { + return nil +} + +func newTestFileStateHandler(t *testing.T) *fileStateHandler { + t.Helper() + + return &fileStateHandler{ + state: make(map[string]FileState), + listenDirs: mapset.NewSet[string](), + collectDirs: mapset.NewSet[string](), + activeTopics: mapset.NewSet[string](), + statePath: filepath.Join(t.TempDir(), "file-state.json"), + } +} + func TestRecursiveDirectoryUpdatesRollBackAfterLaterTraversalError(t *testing.T) { t.Parallel() @@ -81,3 +146,214 @@ func TestRecursiveDirectoryUpdatesRollBackAfterLaterTraversalError(t *testing.T) }) } } + +func TestConcurrentStateAccessAndPersistence(t *testing.T) { + t.Parallel() + + handler := newTestFileStateHandler(t) + baseDir := t.TempDir() + const iterations = 500 + + var wg sync.WaitGroup + errs := make(chan error, iterations) + + wg.Add(4) + go func() { + defer wg.Done() + for i := range iterations { + filename := filepath.Join(baseDir, fmt.Sprintf("%d.log", i)) + handler.setFileState(filename, FileState{Size: int64(i), IsListening: true}) + } + }() + go func() { + defer wg.Done() + for range iterations { + _ = handler.Files() + } + }() + go func() { + defer wg.Done() + for range iterations { + if err := handler.saveState(); err != nil { + errs <- err + } + } + }() + go func() { + defer wg.Done() + for i := range iterations { + _, _ = handler.getFileState(filepath.Join(baseDir, fmt.Sprintf("%d.log", i))) + } + }() + + wg.Wait() + close(errs) + for err := range errs { + require.NoError(t, err) + } +} + +func TestConcurrentDirectoryUpdates(t *testing.T) { + t.Parallel() + + handler := newTestFileStateHandler(t) + listenDir := t.TempDir() + collectDir := t.TempDir() + const iterations = 25 + + var wg sync.WaitGroup + errs := make(chan error, iterations*2) + + wg.Add(2) + go func() { + defer wg.Done() + for range iterations { + if err := handler.UpdateListenDirs(config.DefaultModConfConfig{ + ListenDirs: []string{listenDir}, + }); err != nil { + errs <- err + } + } + }() + go func() { + defer wg.Done() + for range iterations { + if err := handler.UpdateCollectDirs(nil, config.DefaultModConfConfig{ + CollectDirs: []string{collectDir}, + }); err != nil { + errs <- err + } + } + }() + + wg.Wait() + close(errs) + for err := range errs { + require.NoError(t, err) + } +} + +func TestProcessListenFileDoesNotRevertConcurrentProcessedState(t *testing.T) { + t.Parallel() + + handler := newTestFileStateHandler(t) + filePath := filepath.Join(t.TempDir(), "active.log") + require.NoError(t, os.WriteFile(filePath, []byte("log"), 0o600)) + info, err := os.Stat(filePath) + require.NoError(t, err) + + handler.setFileState(filePath, FileState{ + Size: info.Size(), + ModifyTime: info.ModTime().UnixNano(), + IsListening: true, + ProcessState: processStateSeenOnce, + }) + blockingHandler := &blockingFinishedHandler{ + started: make(chan struct{}), + release: make(chan struct{}), + } + handler.handlers = []file_handlers.Interface{blockingHandler} + + done := make(chan struct{}) + go func() { + defer close(done) + handler.processListenFile(filePath, info, 24) + }() + + <-blockingHandler.started + require.NoError(t, handler.MarkProcessedFile(filePath)) + close(blockingHandler.release) + <-done + + state, exists := handler.getFileState(filePath) + require.True(t, exists) + require.Equal(t, processStateProcessed, state.ProcessState) +} + +func TestFailedFileIsNotReadyUntilItsContentsChange(t *testing.T) { + t.Parallel() + + handler := newTestFileStateHandler(t) + filePath := filepath.Join(t.TempDir(), "failed.log") + require.NoError(t, os.WriteFile(filePath, []byte("old"), 0o600)) + info, err := os.Stat(filePath) + require.NoError(t, err) + + handler.setFileState(filePath, FileState{ + Size: info.Size(), + ModifyTime: info.ModTime().UnixNano(), + IsListening: true, + ProcessState: processStateReadyToProcess, + }) + require.NoError(t, handler.MarkFailedFile(filePath)) + require.NoError(t, handler.UpdateFilesProcessState()) + require.NoError(t, handler.UpdateFilesProcessState()) + require.Empty(t, handler.Files(FilterReadyToProcess())) + + require.NoError(t, os.WriteFile(filePath, []byte("new contents"), 0o600)) + modifiedInfo, err := os.Stat(filePath) + require.NoError(t, err) + handler.handlers = []file_handlers.Interface{&blockingFinishedHandler{}} + + handler.processListenFile(filePath, modifiedInfo, 24) + + state, exists := handler.getFileState(filePath) + require.True(t, exists) + require.Equal(t, processStateUnprocessed, state.ProcessState) + require.Equal(t, modifiedInfo.Size(), state.Size) +} + +func TestFailedFileRetriesWhenOnlyModTimeNanosecondsChange(t *testing.T) { + t.Parallel() + + handler := newTestFileStateHandler(t) + filePath := filepath.Join(t.TempDir(), "same-size.log") + require.NoError(t, os.WriteFile(filePath, []byte("same"), 0o600)) + info, err := os.Stat(filePath) + require.NoError(t, err) + + oldModTime := time.Now().Truncate(time.Second).Add(100 * time.Nanosecond) + newModTime := oldModTime.Add(time.Nanosecond) + handler.setFileState(filePath, FileState{ + Size: info.Size(), + ModifyTime: oldModTime.UnixNano(), + IsListening: true, + ProcessState: processStateFailed, + }) + handler.handlers = []file_handlers.Interface{&blockingFinishedHandler{}} + + handler.processListenFile(filePath, fileInfoWithModTime{ + FileInfo: info, + modTime: newModTime, + }, 24) + + state, exists := handler.getFileState(filePath) + require.True(t, exists) + require.Equal(t, processStateUnprocessed, state.ProcessState) + require.Equal(t, newModTime.UnixNano(), state.ModifyTime) +} + +func TestRestoreFileStatesPreservesConcurrentProcessState(t *testing.T) { + t.Parallel() + + handler := newTestFileStateHandler(t) + filePath := filepath.Join(t.TempDir(), "restore.log") + original := FileState{ + Size: 10, + ModifyTime: 20, + IsListening: true, + ProcessState: processStateSeenOnce, + } + handler.setFileState(filePath, original) + + checkpoints := make(map[string]fileStateCheckpoint) + handler.checkpointFileState(checkpoints, filePath) + require.NoError(t, handler.MarkProcessedFile(filePath)) + handler.restoreFileStates(checkpoints) + + state, exists := handler.getFileState(filePath) + require.True(t, exists) + require.Equal(t, processStateProcessed, state.ProcessState) + require.Equal(t, original.Size, state.Size) + require.Equal(t, original.ModifyTime, state.ModifyTime) +} diff --git a/internal/mod/rule/handler.go b/internal/mod/rule/handler.go index 5057c46..ab81130 100644 --- a/internal/mod/rule/handler.go +++ b/internal/mod/rule/handler.go @@ -18,6 +18,7 @@ import ( "bytes" "context" "encoding/json" + stderrors "errors" "os" "path/filepath" "strings" @@ -60,6 +61,7 @@ type queuedRuleMessage struct { item rule_engine.RuleItem payload []byte reservedBytes int64 + result chan error } type CustomRuleHandler struct { @@ -74,6 +76,8 @@ type CustomRuleHandler struct { ruleMessageChan chan queuedRuleMessage engine Engine queuedBytes *queuedbytes.Budget + inFlightMu sync.Mutex + inFlightFiles map[string]struct{} // Master-slave components (optional) slaveRegistry *master.SlaveRegistry @@ -110,6 +114,7 @@ func NewRuleHandler(reqClient api.RequestClient, confManager config.ConfManager, listenChan: make(chan string, 1000), ruleMessageChan: make(chan queuedRuleMessage, 1000), engine: Engine{reqClient: reqClient, ruleDebounceTime: make(map[string]*time.Time)}, + inFlightFiles: make(map[string]struct{}), pubSub: pubSub, queuedBytes: queuedBytes, cleanCollectInfo: func(info model.CollectInfo) string { @@ -140,7 +145,14 @@ func (c *CustomRuleHandler) Run(ctx context.Context) { case <-t.C: configTicker.Reset(config.ReloadRulesInterval) - appConfig := c.confManager.LoadWithRemote() + appConfig, err := c.confManager.LoadWithRemote() + if err != nil { + if appConfig == nil { + log.Errorf("load rule config: %v", err) + continue + } + log.Warnf("Unable to reload rule config, using last-known-good config: %v", err) + } confConfig, ok := appConfig.Mod.Config.(config.DefaultModConfConfig) if ok { *modConfig = confConfig @@ -157,7 +169,11 @@ func (c *CustomRuleHandler) Run(ctx context.Context) { apiRules, err := c.reqClient.ListDeviceDiagnosisRules(device.GetName()) if err != nil { log.Errorf("list device diagnosis rules: %v", err) - c.errChan <- errors.Errorf("list device diagnosis rules: %v", err) + select { + case c.errChan <- errors.Errorf("list device diagnosis rules: %v", err): + case <-ctx.Done(): + return + } continue } log.Infof("received rules: %d", len(apiRules)) @@ -200,7 +216,7 @@ func (c *CustomRuleHandler) Run(ctx context.Context) { case <-t.C: listenTicker.Reset(config.RuleCheckListenFilesInterval) - c.sendFilesToBeProcessed(modConfig) + c.sendFilesToBeProcessed(ctx, modConfig) case <-ctx.Done(): return } @@ -229,8 +245,7 @@ func (c *CustomRuleHandler) Run(ctx context.Context) { case <-t.C: t.Reset(config.RuleScanCollectInfosInterval) - //nolint: contextcheck// context is checked in the parent goroutine - c.scanCollectInfosAndHandle(modConfig) + c.scanCollectInfosAndHandle(ctx, modConfig) case <-ctx.Done(): log.Infof("collect info handler stopped") return @@ -262,7 +277,10 @@ func (c *CustomRuleHandler) handleSubMsg(ctx context.Context) { messages, err := c.pubSub.Subscribe(ctx, constant.TopicRuleMsg) if err != nil { log.Errorf("subscribe to rule message: %v", err) - c.errChan <- errors.Wrap(err, "subscribe to rule message") + select { + case c.errChan <- errors.Wrap(err, "subscribe to rule message"): + case <-ctx.Done(): + } return } @@ -320,13 +338,20 @@ func (c *CustomRuleHandler) handleRuleMessage(ctx context.Context, msg *gcmessag func (c *CustomRuleHandler) consumeRuleMessage(message queuedRuleMessage) { if message.payload != nil { defer c.queuedBytes.Release(message.reservedBytes) - if err := consumeHTTPRuleBatch(message.payload, c.engine.ConsumeNext); err != nil { + if err := consumeHTTPRuleBatch(message.payload, func(item rule_engine.RuleItem) { + if consumeErr := c.engine.ConsumeNext(item); consumeErr != nil { + log.Errorf("consume HTTP rule item: %v", consumeErr) + } + }); err != nil { log.Errorf("consume HTTP rule message batch: %v", err) } return } - c.engine.ConsumeNext(message.item) + consumeErr := c.engine.ConsumeNext(message.item) + if message.result != nil { + message.result <- consumeErr + } } func (c *CustomRuleHandler) releaseQueuedRuleMessages() { @@ -335,6 +360,8 @@ func (c *CustomRuleHandler) releaseQueuedRuleMessages() { case message := <-c.ruleMessageChan: if message.payload != nil { c.queuedBytes.Release(message.reservedBytes) + } else if message.result != nil { + message.result <- context.Canceled } default: return @@ -365,7 +392,10 @@ func consumeHTTPRuleBatch(payload []byte, consume func(rule_engine.RuleItem)) er return err } -func (c *CustomRuleHandler) sendFilesToBeProcessed(modConfig *config.DefaultModConfConfig) { +func (c *CustomRuleHandler) sendFilesToBeProcessed( + ctx context.Context, + modConfig *config.DefaultModConfConfig, +) { if len(modConfig.ListenDirs) == 0 { return } @@ -384,12 +414,44 @@ func (c *CustomRuleHandler) sendFilesToBeProcessed(modConfig *config.DefaultModC file_state_handler.FilterIsListening(), file_state_handler.FilterReadyToProcess(), ) { - err := c.listenFileStateHandler.MarkProcessedFile(fileState.Pathname) - if err != nil { - log.Errorf("mark processed file: %v", err) + if !c.addFileInFlight(fileState.Pathname) { + continue } - c.listenChan <- fileState.Pathname + + select { + case c.listenChan <- fileState.Pathname: + case <-ctx.Done(): + c.releaseFileInFlight(fileState.Pathname) + return + } + } +} + +func (c *CustomRuleHandler) addFileInFlight(filename string) bool { + c.inFlightMu.Lock() + defer c.inFlightMu.Unlock() + + if c.inFlightFiles == nil { + c.inFlightFiles = make(map[string]struct{}) + } + if _, exists := c.inFlightFiles[filename]; exists { + return false } + c.inFlightFiles[filename] = struct{}{} + return true +} + +func (c *CustomRuleHandler) releaseFileInFlight(filename string) { + c.inFlightMu.Lock() + defer c.inFlightMu.Unlock() + delete(c.inFlightFiles, filename) +} + +func (c *CustomRuleHandler) isFileInFlight(filename string) bool { + c.inFlightMu.Lock() + defer c.inFlightMu.Unlock() + _, exists := c.inFlightFiles[filename] + return exists } func (c *CustomRuleHandler) processListenedFilesAndSendMessages( @@ -411,6 +473,7 @@ func (c *CustomRuleHandler) processListenedFilesAndSendMessages( select { case semaphore <- struct{}{}: case <-ctx.Done(): + c.releaseFileInFlight(fileToProcess) wg.Wait() return } @@ -419,8 +482,23 @@ func (c *CustomRuleHandler) processListenedFilesAndSendMessages( go func(filename string) { defer wg.Done() defer func() { <-semaphore }() // release the semaphore + defer c.releaseFileInFlight(filename) + + if err := c.processFileWithRule(ctx, filename); err != nil { + if !shouldMarkFileFailed(ctx, err) { + return + } + log.Errorf("process file %s with rule: %v", filename, err) + if markErr := c.listenFileStateHandler.MarkFailedFile(filename); markErr != nil { + log.Errorf("mark failed file %s: %v", filename, markErr) + } + return + } - c.processFileWithRule(ctx, filename) + if err := c.listenFileStateHandler.MarkProcessedFile(filename); err != nil { + log.Errorf("mark processed file: %v", err) + return + } log.Infof("Finished processing file: %v", filename) }(fileToProcess) case <-ctx.Done(): @@ -430,24 +508,62 @@ func (c *CustomRuleHandler) processListenedFilesAndSendMessages( } } +func shouldMarkFileFailed(ctx context.Context, err error) bool { + if ctx.Err() != nil { + return false + } + return !errors.Is(err, context.Canceled) && + !errors.Is(err, context.DeadlineExceeded) +} + func (c *CustomRuleHandler) processFileWithRule( ctx context.Context, filename string, -) { +) error { log.Infof("RuleEngine exec file: %v", filename) handler := c.listenFileStateHandler.GetFileHandler(filename) if handler == nil { // this should not happen - log.Errorf("get file handler failed for file: %v", filename) - return + return errors.Errorf("get file handler failed for file: %v", filename) } - handler.SendRuleItems(filename, c.engine.activeTopicSet(), func(item rule_engine.RuleItem) bool { - return c.enqueueFileRuleItem(ctx, item) - }) + var consumeErr error + handlerErr := handler.SendRuleItems( + ctx, + filename, + c.engine.activeTopicSet(), + func(item rule_engine.RuleItem) bool { + result := make(chan error, 1) + select { + case c.ruleMessageChan <- queuedRuleMessage{ + item: item, + result: result, + }: + case <-ctx.Done(): + return false + } + + select { + case itemConsumeErr := <-result: + if itemConsumeErr != nil { + consumeErr = stderrors.Join( + consumeErr, + errors.Wrap(itemConsumeErr, "consume rule item"), + ) + } + return true + case <-ctx.Done(): + return false + } + }, + ) + return stderrors.Join(handlerErr, consumeErr) } -func (c *CustomRuleHandler) enqueueFileRuleItem(ctx context.Context, item rule_engine.RuleItem) bool { +func (c *CustomRuleHandler) enqueueFileRuleItem( + ctx context.Context, + item rule_engine.RuleItem, +) bool { select { case c.ruleMessageChan <- queuedRuleMessage{item: item}: return true @@ -457,7 +573,10 @@ func (c *CustomRuleHandler) enqueueFileRuleItem(ctx context.Context, item rule_e } // scanCollectInfosAndHandle handles all collect info files within the collect info dir. -func (c *CustomRuleHandler) scanCollectInfosAndHandle(modConfig *config.DefaultModConfConfig) { +func (c *CustomRuleHandler) scanCollectInfosAndHandle( + ctx context.Context, + modConfig *config.DefaultModConfConfig, +) { log.Infof("Starts to scan collect info dir") // Search for files under the collect info dir and handles them @@ -465,7 +584,10 @@ func (c *CustomRuleHandler) scanCollectInfosAndHandle(modConfig *config.DefaultM entries, err := os.ReadDir(collectInfoDir) if err != nil { log.Errorf("read collect info dir: %v", err) - c.errChan <- errors.Wrap(err, "read collect info dir") + select { + case c.errChan <- errors.Wrap(err, "read collect info dir"): + case <-ctx.Done(): + } return } @@ -486,10 +608,18 @@ func (c *CustomRuleHandler) scanCollectInfosAndHandle(modConfig *config.DefaultM log.Infof("Found %d collect info files", len(collectInfoIds)) for _, collectInfoId := range collectInfoIds { + if ctx.Err() != nil { + return + } + collectInfo := &model.CollectInfo{} if err := collectInfo.Load(collectInfoId); err != nil { log.Errorf("load collect info: %v", err) - c.errChan <- errors.Wrap(err, "load collect info") + select { + case c.errChan <- errors.Wrap(err, "load collect info"): + case <-ctx.Done(): + return + } continue } @@ -507,13 +637,17 @@ func (c *CustomRuleHandler) scanCollectInfosAndHandle(modConfig *config.DefaultM } log.WithField("collectID", collectInfo.Id).Infof("Found collect info to handle") - c.handleCollectInfo(*collectInfo, *modConfig) + c.handleCollectInfo(ctx, *collectInfo, *modConfig) } log.Infof("Finished scanning collect info dir, found %d collect info files", len(collectInfoIds)) } // handleCollectInfo handles a single the collect info. -func (c *CustomRuleHandler) handleCollectInfo(info model.CollectInfo, modConfig config.DefaultModConfConfig) { +func (c *CustomRuleHandler) handleCollectInfo( + ctx context.Context, + info model.CollectInfo, + modConfig config.DefaultModConfConfig, +) { ruleName, _ := info.DiagnosisTask["rule_name"].(string) ruleDisplayName, _ := info.DiagnosisTask["rule_display_name"].(string) collectLog := log.WithFields(log.Fields{ @@ -614,7 +748,7 @@ func (c *CustomRuleHandler) handleCollectInfo(info model.CollectInfo, modConfig // Get slave files if master-slave is enabled var slaveFiles []master.SlaveFileInfo if c.slaveRegistry != nil && c.masterClient != nil && c.masterConfig != nil { - ctx, cancel := context.WithTimeout(context.Background(), c.masterConfig.RequestTimeout) + requestCtx, cancel := context.WithTimeout(ctx, c.masterConfig.RequestTimeout) defer cancel() taskReq := &master.TaskRequest{ @@ -627,7 +761,7 @@ func (c *CustomRuleHandler) handleCollectInfo(info model.CollectInfo, modConfig RecursivelyWalkDirs: modConfig.RecursivelyWalkDirs, } - responses := c.masterClient.RequestAllSlaveFilesByContent(ctx, c.slaveRegistry, taskReq) + responses := c.masterClient.RequestAllSlaveFilesByContent(requestCtx, c.slaveRegistry, taskReq) switch master.TaskResponsesTimeWindowErrorCode(responses) { case master.TaskErrorCodeInvalidTimeWindow: collectLog.Errorf("a slave rejected the collect info time window as permanently invalid, cleaning") diff --git a/internal/mod/rule/handler_test.go b/internal/mod/rule/handler_test.go index 74ab67f..9ea3d1f 100644 --- a/internal/mod/rule/handler_test.go +++ b/internal/mod/rule/handler_test.go @@ -16,6 +16,8 @@ package rule import ( "context" + "errors" + "sync" "testing" "time" @@ -24,6 +26,16 @@ import ( "github.com/coscene-io/coscout/internal/mod/rule/file_handlers" "github.com/coscene-io/coscout/internal/mod/rule/file_state_handler" "github.com/coscene-io/coscout/internal/model" + "github.com/coscene-io/coscout/pkg/rule_engine" + mapset "github.com/deckarep/golang-set/v2" + "github.com/stretchr/testify/require" +) + +var ( + errTestReadFailed = errors.New("read failed") + errTestCancellationSentinel = errors.New("handler lost cancellation cause") + errTestConsumeFailed = errors.New("consume failed") + errTestFirstItemFailed = errors.New("first item failed") ) func TestHandleCollectInfoTimeWindowLifecycle(t *testing.T) { @@ -96,7 +108,7 @@ func TestHandleCollectInfoTimeWindowLifecycle(t *testing.T) { }, } - handler.handleCollectInfo(info, config.DefaultModConfConfig{}) + handler.handleCollectInfo(t.Context(), info, config.DefaultModConfConfig{}) if cleaned != tc.wantClean { t.Fatalf("cleaned = %v, want %v", cleaned, tc.wantClean) @@ -183,6 +195,753 @@ func (*recordingRuleFileStateHandler) MarkProcessedFile(string) error { return nil } +func (*recordingRuleFileStateHandler) MarkFailedFile(string) error { + return nil +} + func (*recordingRuleFileStateHandler) GetFileHandler(string) file_handlers.Interface { return nil } + +type fakeFileStateHandler struct { + files []file_state_handler.FileState + fileHandler file_handlers.Interface + + mu sync.Mutex + processedFile []string + failedFile []string +} + +func (f *fakeFileStateHandler) UpdateListenDirs(config.DefaultModConfConfig) error { + return nil +} + +func (f *fakeFileStateHandler) UpdateCollectDirs([]string, config.DefaultModConfConfig) error { + return nil +} + +func (f *fakeFileStateHandler) Files(...file_state_handler.FileFilter) []file_state_handler.FileState { + return f.files +} + +func (f *fakeFileStateHandler) UpdateFilesProcessState() error { + return nil +} + +func (f *fakeFileStateHandler) MarkProcessedFile(filename string) error { + f.mu.Lock() + defer f.mu.Unlock() + f.processedFile = append(f.processedFile, filename) + return nil +} + +func (f *fakeFileStateHandler) MarkFailedFile(filename string) error { + f.mu.Lock() + defer f.mu.Unlock() + f.failedFile = append(f.failedFile, filename) + return nil +} + +func (f *fakeFileStateHandler) GetFileHandler(string) file_handlers.Interface { + return f.fileHandler +} + +func (f *fakeFileStateHandler) processedCount() int { + f.mu.Lock() + defer f.mu.Unlock() + return len(f.processedFile) +} + +func (f *fakeFileStateHandler) failedCount() int { + f.mu.Lock() + defer f.mu.Unlock() + return len(f.failedFile) +} + +type cancellationBlockingHandler struct { + started chan struct{} + err error +} + +func (h *cancellationBlockingHandler) CheckFilePath(string) bool { + return true +} + +func (h *cancellationBlockingHandler) GetStartTimeEndTime(string) (*time.Time, *time.Time, error) { + return nil, nil, nil +} + +func (h *cancellationBlockingHandler) GetFileSize(string) (int64, error) { + return 0, nil +} + +func (h *cancellationBlockingHandler) IsFinished(string) bool { + return true +} + +func (h *cancellationBlockingHandler) SendRuleItems( + ctx context.Context, + _ string, + _ mapset.Set[string], + _ func(rule_engine.RuleItem) bool, +) error { + if h.started != nil { + select { + case h.started <- struct{}{}: + default: + } + } + <-ctx.Done() + if h.err != nil { + return h.err + } + return ctx.Err() +} + +type resultFileHandler struct { + err error +} + +func (h *resultFileHandler) CheckFilePath(string) bool { + return true +} + +func (h *resultFileHandler) GetStartTimeEndTime(string) (*time.Time, *time.Time, error) { + return nil, nil, nil +} + +func (h *resultFileHandler) GetFileSize(string) (int64, error) { + return 0, nil +} + +func (h *resultFileHandler) IsFinished(string) bool { + return true +} + +func (h *resultFileHandler) SendRuleItems( + context.Context, + string, + mapset.Set[string], + func(rule_engine.RuleItem) bool, +) error { + return h.err +} + +type itemProducingFileHandler struct { + produced chan struct{} +} + +func (h *itemProducingFileHandler) CheckFilePath(string) bool { + return true +} + +func (h *itemProducingFileHandler) GetStartTimeEndTime(string) (*time.Time, *time.Time, error) { + return nil, nil, nil +} + +func (h *itemProducingFileHandler) GetFileSize(string) (int64, error) { + return 0, nil +} + +func (h *itemProducingFileHandler) IsFinished(string) bool { + return true +} + +func (h *itemProducingFileHandler) SendRuleItems( + ctx context.Context, + filename string, + _ mapset.Set[string], + sendRuleItem func(rule_engine.RuleItem) bool, +) error { + if h.produced != nil { + close(h.produced) + } + if sendRuleItem(rule_engine.RuleItem{ + Source: filename, + Topic: "/fault", + }) { + return nil + } + if ctx.Err() != nil { + return ctx.Err() + } + return nil +} + +type multiItemFileHandler struct { + resultFileHandler + done chan struct{} +} + +func (h *multiItemFileHandler) SendRuleItems( + ctx context.Context, + filename string, + _ mapset.Set[string], + sendRuleItem func(rule_engine.RuleItem) bool, +) error { + defer close(h.done) + for index := 1; index <= 2; index++ { + if !sendRuleItem(rule_engine.RuleItem{ + Msg: map[string]interface{}{"index": index}, + Source: filename, + Topic: "/fault", + }) { + if ctx.Err() != nil { + return ctx.Err() + } + return nil + } + } + return nil +} + +func TestSendFilesToBeProcessedCancellationDoesNotMarkUnqueuedFile(t *testing.T) { + t.Parallel() + + stateHandler := &fakeFileStateHandler{ + files: []file_state_handler.FileState{{Pathname: "pending.log"}}, + } + listenChan := make(chan string, 1) + listenChan <- "already-queued.log" + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: listenChan, + } + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + done := make(chan struct{}) + listenDir := t.TempDir() + + go func() { + defer close(done) + handler.sendFilesToBeProcessed(ctx, &config.DefaultModConfConfig{ + ListenDirs: []string{listenDir}, + }) + }() + + select { + case <-done: + t.Fatal("sendFilesToBeProcessed returned before the full channel was cancelled") + case <-time.After(50 * time.Millisecond): + } + require.Zero(t, stateHandler.processedCount()) + + cancel() + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) + require.Zero(t, stateHandler.processedCount()) + require.Zero(t, stateHandler.failedCount()) + require.False(t, handler.isFileInFlight("pending.log")) +} + +func TestProcessListenedFilesCancellationInterruptsSemaphoreWait(t *testing.T) { + t.Parallel() + + stateHandler := &fakeFileStateHandler{ + fileHandler: &cancellationBlockingHandler{}, + } + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: make(chan string, 2), + ruleMessageChan: make(chan queuedRuleMessage), + } + handler.listenChan <- "first.log" + handler.listenChan <- "second.log" + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + done := make(chan struct{}) + go func() { + defer close(done) + handler.processListenedFilesAndSendMessages(ctx, 1) + }() + + time.Sleep(50 * time.Millisecond) + cancel() + + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) +} + +func TestCancelledFileProcessingDoesNotMarkFileProcessed(t *testing.T) { + t.Parallel() + + started := make(chan struct{}, 1) + stateHandler := &fakeFileStateHandler{ + files: []file_state_handler.FileState{{Pathname: "pending.log"}}, + fileHandler: &cancellationBlockingHandler{started: started}, + } + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: make(chan string, 1), + ruleMessageChan: make(chan queuedRuleMessage), + } + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + handler.sendFilesToBeProcessed(ctx, &config.DefaultModConfConfig{ + ListenDirs: []string{t.TempDir()}, + }) + require.Zero(t, stateHandler.processedCount()) + require.True(t, handler.isFileInFlight("pending.log")) + + done := make(chan struct{}) + go func() { + defer close(done) + handler.processListenedFilesAndSendMessages(ctx, 1) + }() + require.Eventually(t, func() bool { + select { + case <-started: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) + + cancel() + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) + require.Zero(t, stateHandler.processedCount()) + require.Zero(t, stateHandler.failedCount()) + require.False(t, handler.isFileInFlight("pending.log")) +} + +func TestCancellationDoesNotMarkFailedWhenHandlerReturnsSentinelError(t *testing.T) { + t.Parallel() + + started := make(chan struct{}, 1) + stateHandler := &fakeFileStateHandler{ + files: []file_state_handler.FileState{{Pathname: "cancelled-sentinel.log"}}, + fileHandler: &cancellationBlockingHandler{ + started: started, + err: errTestCancellationSentinel, + }, + } + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: make(chan string, 1), + ruleMessageChan: make(chan queuedRuleMessage), + } + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + handler.sendFilesToBeProcessed(ctx, &config.DefaultModConfConfig{ + ListenDirs: []string{t.TempDir()}, + }) + + done := make(chan struct{}) + go func() { + defer close(done) + handler.processListenedFilesAndSendMessages(ctx, 1) + }() + require.Eventually(t, func() bool { + select { + case <-started: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) + + cancel() + require.False(t, shouldMarkFileFailed(ctx, errTestCancellationSentinel)) + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) + require.Zero(t, stateHandler.processedCount()) + require.Zero(t, stateHandler.failedCount()) + require.False(t, handler.isFileInFlight("cancelled-sentinel.log")) +} + +func TestCompletedFileProcessingMarksFileProcessed(t *testing.T) { + t.Parallel() + + stateHandler := &fakeFileStateHandler{ + files: []file_state_handler.FileState{{Pathname: "completed.log"}}, + fileHandler: &resultFileHandler{}, + } + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: make(chan string, 1), + ruleMessageChan: make(chan queuedRuleMessage), + } + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + handler.sendFilesToBeProcessed(ctx, &config.DefaultModConfConfig{ + ListenDirs: []string{t.TempDir()}, + }) + require.Zero(t, stateHandler.processedCount()) + + done := make(chan struct{}) + go func() { + defer close(done) + handler.processListenedFilesAndSendMessages(ctx, 1) + }() + require.Eventually(t, func() bool { + return stateHandler.processedCount() == 1 + }, time.Second, 10*time.Millisecond) + require.Zero(t, stateHandler.failedCount()) + require.False(t, handler.isFileInFlight("completed.log")) + + cancel() + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) +} + +func TestFailedFileProcessingDoesNotMarkFileProcessed(t *testing.T) { + t.Parallel() + + stateHandler := &fakeFileStateHandler{ + files: []file_state_handler.FileState{{Pathname: "failed.log"}}, + fileHandler: &resultFileHandler{ + err: errTestReadFailed, + }, + } + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: make(chan string, 1), + ruleMessageChan: make(chan queuedRuleMessage), + } + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + handler.sendFilesToBeProcessed(ctx, &config.DefaultModConfConfig{ + ListenDirs: []string{t.TempDir()}, + }) + require.True(t, handler.isFileInFlight("failed.log")) + + done := make(chan struct{}) + go func() { + defer close(done) + handler.processListenedFilesAndSendMessages(ctx, 1) + }() + require.Eventually(t, func() bool { + return stateHandler.failedCount() == 1 && + !handler.isFileInFlight("failed.log") + }, time.Second, 10*time.Millisecond) + + require.Zero(t, stateHandler.processedCount()) + cancel() + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) +} + +func TestSendFilesToBeProcessedDoesNotEnqueueInFlightFileTwice(t *testing.T) { + t.Parallel() + + stateHandler := &fakeFileStateHandler{ + files: []file_state_handler.FileState{{Pathname: "pending.log"}}, + } + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: make(chan string, 2), + } + ctx := t.Context() + modConfig := &config.DefaultModConfConfig{ + ListenDirs: []string{t.TempDir()}, + } + + handler.sendFilesToBeProcessed(ctx, modConfig) + handler.sendFilesToBeProcessed(ctx, modConfig) + + require.Len(t, handler.listenChan, 1) + require.Zero(t, stateHandler.processedCount()) + require.True(t, handler.isFileInFlight("pending.log")) +} + +func TestProducedRuleItemWithoutConsumerAckIsNotMarkedProcessedOnCancel(t *testing.T) { + t.Parallel() + + produced := make(chan struct{}) + stateHandler := &fakeFileStateHandler{ + files: []file_state_handler.FileState{{Pathname: "pending-ack.log"}}, + fileHandler: &itemProducingFileHandler{ + produced: produced, + }, + } + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: make(chan string, 1), + ruleMessageChan: make(chan queuedRuleMessage, 1), + } + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + handler.sendFilesToBeProcessed(ctx, &config.DefaultModConfConfig{ + ListenDirs: []string{t.TempDir()}, + }) + + done := make(chan struct{}) + go func() { + defer close(done) + handler.processListenedFilesAndSendMessages(ctx, 1) + }() + require.Eventually(t, func() bool { + select { + case <-produced: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) + + var envelope queuedRuleMessage + select { + case envelope = <-handler.ruleMessageChan: + case <-time.After(time.Second): + t.Fatal("rule item was not forwarded to the consumer channel") + } + require.NotNil(t, envelope.result) + require.Zero(t, stateHandler.processedCount()) + + cancel() + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) + require.Zero(t, stateHandler.processedCount()) + require.Zero(t, stateHandler.failedCount()) + select { + case envelope.result <- nil: + case <-time.After(time.Second): + t.Fatal("consumer result handshake blocked after worker cancellation") + } +} + +func TestProducedRuleItemIsMarkedProcessedOnlyAfterConsumerAck(t *testing.T) { + t.Parallel() + + stateHandler := &fakeFileStateHandler{ + files: []file_state_handler.FileState{{Pathname: "acked.log"}}, + fileHandler: &itemProducingFileHandler{}, + } + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: make(chan string, 1), + ruleMessageChan: make(chan queuedRuleMessage, 1), + } + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + handler.sendFilesToBeProcessed(ctx, &config.DefaultModConfConfig{ + ListenDirs: []string{t.TempDir()}, + }) + + done := make(chan struct{}) + go func() { + defer close(done) + handler.processListenedFilesAndSendMessages(ctx, 1) + }() + + var envelope queuedRuleMessage + select { + case envelope = <-handler.ruleMessageChan: + case <-time.After(time.Second): + t.Fatal("rule item was not forwarded to the consumer channel") + } + require.NotNil(t, envelope.result) + require.Zero(t, stateHandler.processedCount()) + + envelope.result <- nil + require.Eventually(t, func() bool { + return stateHandler.processedCount() == 1 + }, time.Second, 10*time.Millisecond) + require.Zero(t, stateHandler.failedCount()) + + cancel() + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) +} + +func TestRuleItemConsumptionErrorMarksFileFailed(t *testing.T) { + t.Parallel() + + action, err := rule_engine.NewAction( + "failing-action", + map[string]interface{}{}, + func(map[string]interface{}) error { + return errTestConsumeFailed + }, + ) + require.NoError(t, err) + stateHandler := &fakeFileStateHandler{ + files: []file_state_handler.FileState{{Pathname: "consume-error.log"}}, + fileHandler: &itemProducingFileHandler{}, + } + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: make(chan string, 1), + ruleMessageChan: make(chan queuedRuleMessage, 1), + engine: Engine{ + rules: []*rule_engine.Rule{testRuntimeRule(t, action)}, + ruleDebounceTime: make(map[string]*time.Time), + publishCollectInfoFn: func(string) error { + return nil + }, + }, + } + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + handler.sendFilesToBeProcessed(ctx, &config.DefaultModConfConfig{ + ListenDirs: []string{t.TempDir()}, + }) + + done := make(chan struct{}) + go func() { + defer close(done) + handler.processListenedFilesAndSendMessages(ctx, 1) + }() + consumerDone := make(chan struct{}) + go func() { + defer close(consumerDone) + select { + case envelope := <-handler.ruleMessageChan: + envelope.result <- handler.engine.ConsumeNext(envelope.item) + case <-ctx.Done(): + } + }() + + require.Eventually(t, func() bool { + return stateHandler.failedCount() == 1 + }, time.Second, 10*time.Millisecond) + require.Zero(t, stateHandler.processedCount()) + require.Eventually(t, func() bool { + select { + case <-consumerDone: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) + + cancel() + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) +} + +func TestRuleItemConsumptionErrorDrainsRemainingHandlerItems(t *testing.T) { + t.Parallel() + + producerDone := make(chan struct{}) + stateHandler := &fakeFileStateHandler{ + files: []file_state_handler.FileState{{Pathname: "multi-item.log"}}, + fileHandler: &multiItemFileHandler{ + done: producerDone, + }, + } + handler := &CustomRuleHandler{ + listenFileStateHandler: stateHandler, + listenChan: make(chan string, 1), + ruleMessageChan: make(chan queuedRuleMessage, 1), + } + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + handler.sendFilesToBeProcessed(ctx, &config.DefaultModConfConfig{ + ListenDirs: []string{t.TempDir()}, + }) + + workerDone := make(chan struct{}) + go func() { + defer close(workerDone) + handler.processListenedFilesAndSendMessages(ctx, 1) + }() + + var consumedIndexes []int + consumerDone := make(chan struct{}) + go func() { + defer close(consumerDone) + for index := range 2 { + select { + case envelope := <-handler.ruleMessageChan: + itemIndex, _ := envelope.item.Msg["index"].(int) + consumedIndexes = append(consumedIndexes, itemIndex) + if index == 0 { + envelope.result <- errTestFirstItemFailed + } else { + envelope.result <- nil + } + case <-ctx.Done(): + return + } + } + }() + + require.Eventually(t, func() bool { + select { + case <-producerDone: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) + require.Eventually(t, func() bool { + select { + case <-consumerDone: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) + require.Equal(t, []int{1, 2}, consumedIndexes) + require.Eventually(t, func() bool { + return stateHandler.failedCount() == 1 + }, time.Second, 10*time.Millisecond) + require.Zero(t, stateHandler.processedCount()) + + cancel() + require.Eventually(t, func() bool { + select { + case <-workerDone: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) +} diff --git a/internal/slave/server_test.go b/internal/slave/server_test.go index 374804f..fb968e5 100644 --- a/internal/slave/server_test.go +++ b/internal/slave/server_test.go @@ -50,6 +50,8 @@ func (updateCollectDirsErrorHandler) UpdateFilesProcessState() error { return ni func (updateCollectDirsErrorHandler) MarkProcessedFile(string) error { return nil } +func (updateCollectDirsErrorHandler) MarkFailedFile(string) error { return nil } + func (updateCollectDirsErrorHandler) GetFileHandler(string) file_handlers.Interface { return nil } type deadlineResponseWriter struct { diff --git a/pkg/mcap_ros2/cdr_reader.go b/pkg/mcap_ros2/cdr_reader.go index d477df3..3c8f5d4 100644 --- a/pkg/mcap_ros2/cdr_reader.go +++ b/pkg/mcap_ros2/cdr_reader.go @@ -22,6 +22,25 @@ import ( "github.com/pkg/errors" ) +const ( + // maxDecodedArrayBytes caps the estimated backing storage of any decoded + // array. It is deliberately independent of the encoded payload size because + // []string and []interface{} amplify small encoded elements into two-word Go + // values. + maxDecodedArrayBytes = 64 << 20 + + // Interfaces and strings both occupy two machine words on the supported + // 64-bit targets. Keeping the conservative 16-byte estimate on 32-bit + // targets only makes the decoder limit stricter. + resultInterfaceSlotBytes = 16 + resultStringHeaderBytes = 16 + + // A decoded complex element owns both its interface slot and at least one + // map produced by readComplexType. Budget conservatively for that object + // graph rather than only for the interface backing array. + resultComplexElementBytes = 128 +) + // CdrReader parses values from CDR data. type CdrReader struct { data []byte @@ -52,6 +71,14 @@ func (r *CdrReader) ByteLength() int { return len(r.data) } +// RemainingBytes returns the number of unread bytes in the CDR payload. +func (r *CdrReader) RemainingBytes() int { + if r.offset >= len(r.data) { + return 0 + } + return len(r.data) - r.offset +} + // Boolean reads an 8-bit value and interprets it as a boolean. func (r *CdrReader) Boolean() (bool, error) { val, err := r.Uint8() @@ -61,13 +88,19 @@ func (r *CdrReader) Boolean() (bool, error) { // Int8 reads a signed 8-bit integer. func (r *CdrReader) Int8() (int8, error) { val, err := r.read(1) - return int8(val[0]), err + if err != nil { + return 0, err + } + return int8(val[0]), nil } // Uint8 reads an unsigned 8-bit integer. func (r *CdrReader) Uint8() (uint8, error) { val, err := r.read(1) - return val[0], err + if err != nil { + return 0, err + } + return val[0], nil } // uint16ToInt16 converts a uint16 to a int16. @@ -79,13 +112,19 @@ func uint16ToInt16(u uint16) int16 { // Int16 reads a signed 16-bit integer. func (r *CdrReader) Int16() (int16, error) { val, err := r.read(2) - return uint16ToInt16(r.byteOrder().Uint16(val)), err + if err != nil { + return 0, err + } + return uint16ToInt16(r.byteOrder().Uint16(val)), nil } // Uint16 reads an unsigned 16-bit integer. func (r *CdrReader) Uint16() (uint16, error) { val, err := r.read(2) - return r.byteOrder().Uint16(val), err + if err != nil { + return 0, err + } + return r.byteOrder().Uint16(val), nil } // uint32ToInt32 converts a uint32 to a int32. @@ -97,13 +136,19 @@ func uint32ToInt32(u uint32) int32 { // Int32 reads a signed 32-bit integer. func (r *CdrReader) Int32() (int32, error) { val, err := r.read(4) - return uint32ToInt32(r.byteOrder().Uint32(val)), err + if err != nil { + return 0, err + } + return uint32ToInt32(r.byteOrder().Uint32(val)), nil } // Uint32 reads an unsigned 32-bit integer. func (r *CdrReader) Uint32() (uint32, error) { val, err := r.read(4) - return r.byteOrder().Uint32(val), err + if err != nil { + return 0, err + } + return r.byteOrder().Uint32(val), nil } // uint64ToInt64 converts a uint64 to a int64. @@ -115,25 +160,37 @@ func uint64ToInt64(u uint64) int64 { // Int64 reads a signed 64-bit integer. func (r *CdrReader) Int64() (int64, error) { val, err := r.read(8) - return uint64ToInt64(r.byteOrder().Uint64(val)), err + if err != nil { + return 0, err + } + return uint64ToInt64(r.byteOrder().Uint64(val)), nil } // Uint64 reads an unsigned 64-bit integer. func (r *CdrReader) Uint64() (uint64, error) { val, err := r.read(8) - return r.byteOrder().Uint64(val), err + if err != nil { + return 0, err + } + return r.byteOrder().Uint64(val), nil } // Float32 reads a 32-bit floating point number. func (r *CdrReader) Float32() (JSONFloat32, error) { val, err := r.read(4) - return JSONFloat32(math.Float32frombits(r.byteOrder().Uint32(val))), err + if err != nil { + return 0, err + } + return JSONFloat32(math.Float32frombits(r.byteOrder().Uint32(val))), nil } // Float64 reads a 64-bit floating point number. func (r *CdrReader) Float64() (JSONFloat64, error) { val, err := r.read(8) - return JSONFloat64(math.Float64frombits(r.byteOrder().Uint64(val))), err + if err != nil { + return 0, err + } + return JSONFloat64(math.Float64frombits(r.byteOrder().Uint64(val))), nil } // String reads a string prefixed with its 32-bit length. @@ -143,17 +200,21 @@ func (r *CdrReader) String() (string, error) { return "", err } - if length <= 1 { - r.offset += int(length) - return "", nil + if uint64(length) > uint64(math.MaxInt) { + return "", errors.Errorf("string declared length %d cannot fit in int", length) } - str, err := r.StringRaw(int(length - 1)) + // The conversion is safe because the platform int bound is checked above. + // #nosec G115 + data, err := r.Uint8Array(int(length)) if err != nil { return "", err } - r.offset++ // Skip null terminator - return str, nil + if length <= 1 { + return "", nil + } + + return string(data[:len(data)-1]), nil } // StringRaw reads a string of the given length. @@ -188,8 +249,8 @@ func (r *CdrReader) align(size int) { // read reads bytes from the current offset and advances the offset. func (r *CdrReader) read(size int) ([]byte, error) { - if r.offset+size > len(r.data) { - return nil, errors.Errorf("attempt to read past end of data (offset: %d, size: %d, data length: %d)", r.offset, size, len(r.data)) + if size < 0 { + return nil, errors.Errorf("attempt to read a negative size %d", size) } // Align before reading if size > 1 @@ -197,13 +258,83 @@ func (r *CdrReader) read(size int) ([]byte, error) { r.align(size) } + if r.offset > len(r.data) || size > len(r.data)-r.offset { + return nil, errors.Errorf("attempt to read past end of data (offset: %d, size: %d, data length: %d)", r.offset, size, len(r.data)) + } + data := r.data[r.offset : r.offset+size] r.offset += size return data, nil } +// validateArrayLength validates a declared array length before it is converted +// to int or used by make/slicing. +func validateArrayLength( + length uint64, + remainingBytes, minEncodedBytes, resultElementBytes int, +) (int, error) { + if length > uint64(math.MaxInt) { + return 0, errors.Errorf("array declared length %d cannot fit in int", length) + } + + // The conversion is safe because the platform int bound is checked above. + // #nosec G115 + intLength := int(length) + return intLength, validateIntArrayLength( + intLength, + remainingBytes, + minEncodedBytes, + resultElementBytes, + ) +} + +func validateIntArrayLength(length, remainingBytes, minEncodedBytes, resultElementBytes int) error { + if length < 0 { + return errors.Errorf("array declared a negative length %d", length) + } + if remainingBytes < 0 || minEncodedBytes <= 0 || resultElementBytes <= 0 { + return errors.Errorf( + "invalid array validation parameters (remaining=%d, minimum encoded element=%d, result element=%d)", + remainingBytes, + minEncodedBytes, + resultElementBytes, + ) + } + + remainingLimit := remainingBytes / minEncodedBytes + if length > remainingLimit { + return errors.Errorf( + "array declared length %d exceeds remaining-data limit %d (remaining=%d bytes, minimum encoded element=%d bytes)", + length, + remainingLimit, + remainingBytes, + minEncodedBytes, + ) + } + + allocationLimit := maxDecodedArrayBytes / resultElementBytes + if length > allocationLimit { + return errors.Errorf( + "array declared length %d exceeds decoder allocation limit %d elements (%d bytes, result element=%d bytes)", + length, + allocationLimit, + maxDecodedArrayBytes, + resultElementBytes, + ) + } + + return nil +} + +func (r *CdrReader) validateArrayLength(length, minEncodedBytes, resultElementBytes int) error { + return validateIntArrayLength(length, r.RemainingBytes(), minEncodedBytes, resultElementBytes) +} + // BooleanArray reads an array of booleans. func (r *CdrReader) BooleanArray(length int) ([]bool, error) { + if err := r.validateArrayLength(length, 1, 1); err != nil { + return nil, err + } result := make([]bool, length) for i := range length { val, err := r.Boolean() @@ -217,6 +348,9 @@ func (r *CdrReader) BooleanArray(length int) ([]bool, error) { // Uint8Array reads an array of uint8 values. func (r *CdrReader) Uint8Array(length int) ([]uint8, error) { + if err := r.validateArrayLength(length, 1, 1); err != nil { + return nil, err + } data := r.data[r.offset : r.offset+length] r.offset += length return data, nil @@ -224,6 +358,9 @@ func (r *CdrReader) Uint8Array(length int) ([]uint8, error) { // Int8Array reads an array of int8 values. func (r *CdrReader) Int8Array(length int) ([]int8, error) { + if err := r.validateArrayLength(length, 1, 1); err != nil { + return nil, err + } result := make([]int8, length) for i := range length { val, err := r.Int8() @@ -237,6 +374,9 @@ func (r *CdrReader) Int8Array(length int) ([]int8, error) { // Int16Array reads an array of int16 values. func (r *CdrReader) Int16Array(length int) ([]int16, error) { + if err := r.validateArrayLength(length, 2, 2); err != nil { + return nil, err + } result := make([]int16, length) for i := range length { val, err := r.Int16() @@ -250,6 +390,9 @@ func (r *CdrReader) Int16Array(length int) ([]int16, error) { // Uint16Array reads an array of uint16 values. func (r *CdrReader) Uint16Array(length int) ([]uint16, error) { + if err := r.validateArrayLength(length, 2, 2); err != nil { + return nil, err + } result := make([]uint16, length) for i := range length { val, err := r.Uint16() @@ -263,6 +406,9 @@ func (r *CdrReader) Uint16Array(length int) ([]uint16, error) { // Int32Array reads an array of int32 values. func (r *CdrReader) Int32Array(length int) ([]int32, error) { + if err := r.validateArrayLength(length, 4, 4); err != nil { + return nil, err + } result := make([]int32, length) for i := range length { val, err := r.Int32() @@ -276,6 +422,9 @@ func (r *CdrReader) Int32Array(length int) ([]int32, error) { // Uint32Array reads an array of uint32 values. func (r *CdrReader) Uint32Array(length int) ([]uint32, error) { + if err := r.validateArrayLength(length, 4, 4); err != nil { + return nil, err + } result := make([]uint32, length) for i := range length { val, err := r.Uint32() @@ -289,6 +438,9 @@ func (r *CdrReader) Uint32Array(length int) ([]uint32, error) { // Int64Array reads an array of int64 values. func (r *CdrReader) Int64Array(length int) ([]int64, error) { + if err := r.validateArrayLength(length, 8, 8); err != nil { + return nil, err + } result := make([]int64, length) for i := range length { val, err := r.Int64() @@ -302,6 +454,9 @@ func (r *CdrReader) Int64Array(length int) ([]int64, error) { // Uint64Array reads an array of uint64 values. func (r *CdrReader) Uint64Array(length int) ([]uint64, error) { + if err := r.validateArrayLength(length, 8, 8); err != nil { + return nil, err + } result := make([]uint64, length) for i := range length { val, err := r.Uint64() @@ -315,6 +470,9 @@ func (r *CdrReader) Uint64Array(length int) ([]uint64, error) { // Float32Array reads an array of float32 values. func (r *CdrReader) Float32Array(length int) ([]JSONFloat32, error) { + if err := r.validateArrayLength(length, 4, 4); err != nil { + return nil, err + } result := make([]JSONFloat32, length) for i := range length { val, err := r.Float32() @@ -328,6 +486,9 @@ func (r *CdrReader) Float32Array(length int) ([]JSONFloat32, error) { // Float64Array reads an array of float64 values. func (r *CdrReader) Float64Array(length int) ([]JSONFloat64, error) { + if err := r.validateArrayLength(length, 8, 8); err != nil { + return nil, err + } result := make([]JSONFloat64, length) for i := range length { val, err := r.Float64() @@ -341,6 +502,9 @@ func (r *CdrReader) Float64Array(length int) ([]JSONFloat64, error) { // StringArray reads an array of strings. func (r *CdrReader) StringArray(length int) ([]string, error) { + if err := r.validateArrayLength(length, 4, resultStringHeaderBytes); err != nil { + return nil, err + } result := make([]string, length) for i := range length { val, err := r.String() diff --git a/pkg/mcap_ros2/message_reader.go b/pkg/mcap_ros2/message_reader.go index 3a45dd1..e44523f 100644 --- a/pkg/mcap_ros2/message_reader.go +++ b/pkg/mcap_ros2/message_reader.go @@ -400,26 +400,22 @@ func readField( return readComplexType(nestedDef, msgDefs, reader) } - var arrayLength int - if ftype.ArraySize != nil { - // Fixed size array - arrayLength = *ftype.ArraySize - } else { - // Dynamic array - read length prefix - length, err := reader.Uint32() - if err != nil { - return nil, err - } - arrayLength = int(length) - - // Handle upper bound if specified - if ftype.IsUpperBound && ftype.ArraySize != nil && arrayLength > *ftype.ArraySize { - arrayLength = *ftype.ArraySize - } + arrayLength, err := readArrayLength(ftype, reader) + if err != nil { + return nil, err + } + length, err := validateArrayLength( + arrayLength, + reader.RemainingBytes(), + 1, + resultComplexElementBytes, + ) + if err != nil { + return nil, err } - array := make([]interface{}, arrayLength) - for i := range arrayLength { + array := make([]interface{}, length) + for i := range length { value, err := readComplexType(nestedDef, msgDefs, reader) if err != nil { return nil, err @@ -431,54 +427,98 @@ func readField( // readPrimitiveArray reads an array of primitive values. func readPrimitiveArray(ftype Type, reader *CdrReader) (interface{}, error) { - var arrayLength int - if ftype.ArraySize != nil { - // Fixed size array - arrayLength = *ftype.ArraySize - } else { - // Dynamic array - read length prefix - length, err := reader.Uint32() - if err != nil { - return nil, err - } - arrayLength = int(length) - - // Handle upper bound if specified - if ftype.IsUpperBound && ftype.ArraySize != nil && arrayLength > *ftype.ArraySize { - arrayLength = *ftype.ArraySize - } + arrayLength, err := readArrayLength(ftype, reader) + if err != nil { + return nil, err + } + minEncodedBytes, resultElementBytes, err := primitiveArrayElementSizes(ftype.Type) + if err != nil { + return nil, err + } + length, err := validateArrayLength( + arrayLength, + reader.RemainingBytes(), + minEncodedBytes, + resultElementBytes, + ) + if err != nil { + return nil, err } switch ftype.Type { case typeBool: - return reader.BooleanArray(arrayLength) + return reader.BooleanArray(length) case typeByte, typeUint8: - return reader.Uint8Array(arrayLength) + return reader.Uint8Array(length) case typeChar, typeInt8: - return reader.Int8Array(arrayLength) + return reader.Int8Array(length) case typeInt16: - return reader.Int16Array(arrayLength) + return reader.Int16Array(length) case typeUint16: - return reader.Uint16Array(arrayLength) + return reader.Uint16Array(length) case typeInt32: - return reader.Int32Array(arrayLength) + return reader.Int32Array(length) case typeUint32: - return reader.Uint32Array(arrayLength) + return reader.Uint32Array(length) case typeInt64: - return reader.Int64Array(arrayLength) + return reader.Int64Array(length) case typeUint64: - return reader.Uint64Array(arrayLength) + return reader.Uint64Array(length) case typeFloat32: - return reader.Float32Array(arrayLength) + return reader.Float32Array(length) case typeFloat64: - return reader.Float64Array(arrayLength) + return reader.Float64Array(length) case typeString: - return reader.StringArray(arrayLength) + return reader.StringArray(length) default: return nil, errors.Errorf("unsupported array type: %s", ftype.Type) } } +func readArrayLength(ftype Type, reader *CdrReader) (uint64, error) { + if ftype.ArraySize != nil && !ftype.IsUpperBound { + if *ftype.ArraySize < 0 { + return 0, errors.Errorf("array declared a negative fixed length %d", *ftype.ArraySize) + } + return uint64(*ftype.ArraySize), nil + } + + length, err := reader.SequenceLength() + if err != nil { + return 0, err + } + if ftype.IsUpperBound && ftype.ArraySize != nil { + if *ftype.ArraySize < 0 { + return 0, errors.Errorf("array declared a negative schema upper bound %d", *ftype.ArraySize) + } + if uint64(length) > uint64(*ftype.ArraySize) { + return 0, errors.Errorf( + "array declared length %d exceeds schema upper bound %d", + length, + *ftype.ArraySize, + ) + } + } + return uint64(length), nil +} + +func primitiveArrayElementSizes(typeName string) (int, int, error) { + switch typeName { + case typeBool, typeByte, typeUint8, typeChar, typeInt8: + return 1, 1, nil + case typeInt16, typeUint16: + return 2, 2, nil + case typeInt32, typeUint32, typeFloat32: + return 4, 4, nil + case typeInt64, typeUint64, typeFloat64: + return 8, 8, nil + case typeString: + return 4, resultStringHeaderBytes, nil + default: + return 0, 0, errors.Errorf("unsupported array type: %s", typeName) + } +} + // readPrimitiveValue reads a single primitive value. func readPrimitiveValue(typeName string, reader *CdrReader) (interface{}, error) { switch typeName { diff --git a/pkg/mcap_ros2/message_reader_test.go b/pkg/mcap_ros2/message_reader_test.go new file mode 100644 index 0000000..4d59ac4 --- /dev/null +++ b/pkg/mcap_ros2/message_reader_test.go @@ -0,0 +1,299 @@ +// Copyright 2026 coScene +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package mcap_ros2 + +import ( + "encoding/binary" + "math" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func dynamicArrayData(length uint32, payload ...byte) []byte { + data := make([]byte, 8, 8+len(payload)) + data[1] = 1 // Little-endian CDR. + binary.LittleEndian.PutUint32(data[4:], length) + return append(data, payload...) +} + +func TestReadPrimitiveArrayRejectsLengthExceedingRemainingData(t *testing.T) { + t.Parallel() + + reader, err := NewCdrReader(dynamicArrayData(math.MaxUint32)) + require.NoError(t, err) + + require.NotPanics(t, func() { + _, err = readPrimitiveArray(Type{Type: typeUint8, IsArray: true}, reader) + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "declared length") + assert.Contains(t, err.Error(), "remaining") +} + +func TestReadComplexArrayRejectsLengthExceedingRemainingData(t *testing.T) { + t.Parallel() + + reader, err := NewCdrReader(dynamicArrayData(math.MaxUint32)) + require.NoError(t, err) + field := Field{ + Type: Type{ + PkgName: "test", + Type: "Empty", + IsArray: true, + }, + Name: "items", + } + msgDefs := map[string]MessageSpecification{ + "test/Empty": {PkgName: "test", MsgName: "Empty"}, + } + + require.NotPanics(t, func() { + _, err = readField(field, msgDefs, reader) + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "declared length") + assert.Contains(t, err.Error(), "remaining") +} + +func TestReadComplexArrayRejectsHardAllocationLimit(t *testing.T) { + t.Parallel() + + length := maxDecodedArrayBytes/resultComplexElementBytes + 1 + reader, err := NewCdrReader(dynamicArrayData(uint32(length), make([]byte, length)...)) + require.NoError(t, err) + field := Field{ + Type: Type{ + PkgName: "test", + Type: "Empty", + IsArray: true, + }, + Name: "items", + } + msgDefs := map[string]MessageSpecification{ + "test/Empty": {PkgName: "test", MsgName: "Empty"}, + } + + _, err = readField(field, msgDefs, reader) + + require.Error(t, err) + assert.Contains(t, err.Error(), "decoder allocation limit") +} + +func TestValidateArrayLengthRejectsHardAllocationLimit(t *testing.T) { + t.Parallel() + + length := uint64(maxDecodedArrayBytes/resultInterfaceSlotBytes + 1) + + _, err := validateArrayLength(length, math.MaxInt, 1, resultInterfaceSlotBytes) + + require.Error(t, err) + assert.Contains(t, err.Error(), "decoder allocation limit") +} + +func TestReadPrimitiveArrayDecodesValidDynamicArray(t *testing.T) { + t.Parallel() + + reader, err := NewCdrReader(dynamicArrayData(3, 1, 2, 3)) + require.NoError(t, err) + + value, err := readPrimitiveArray(Type{Type: typeUint8, IsArray: true}, reader) + + require.NoError(t, err) + assert.Equal(t, []uint8{1, 2, 3}, value) +} + +func TestReadComplexArrayDecodesValidDynamicArray(t *testing.T) { + t.Parallel() + + reader, err := NewCdrReader(dynamicArrayData(2, 7, 8)) + require.NoError(t, err) + field := Field{ + Type: Type{ + PkgName: "test", + Type: "Empty", + IsArray: true, + }, + Name: "items", + } + msgDefs := map[string]MessageSpecification{ + "test/Empty": {PkgName: "test", MsgName: "Empty"}, + } + + value, err := readField(field, msgDefs, reader) + + require.NoError(t, err) + assert.Equal(t, []interface{}{ + map[string]interface{}{}, + map[string]interface{}{}, + }, value) +} + +func TestReadPrimitiveArrayDecodesBoundedDynamicArray(t *testing.T) { + t.Parallel() + + bound := 3 + reader, err := NewCdrReader(dynamicArrayData(2, 7, 8)) + require.NoError(t, err) + + value, err := readPrimitiveArray(Type{ + Type: typeUint8, + IsArray: true, + ArraySize: &bound, + IsUpperBound: true, + }, reader) + + require.NoError(t, err) + assert.Equal(t, []uint8{7, 8}, value) +} + +func TestReadPrimitiveArrayRejectsSchemaBoundViolation(t *testing.T) { + t.Parallel() + + bound := 3 + reader, err := NewCdrReader(dynamicArrayData(4, 1, 2, 3, 4)) + require.NoError(t, err) + + _, err = readPrimitiveArray(Type{ + Type: typeUint8, + IsArray: true, + ArraySize: &bound, + IsUpperBound: true, + }, reader) + + require.Error(t, err) + assert.Contains(t, err.Error(), "schema upper bound") +} + +func TestReadPrimitiveFixedArrayRejectsTruncatedDataWithoutPanic(t *testing.T) { + t.Parallel() + + arraySize := 3 + reader, err := NewCdrReader([]byte{0, 1, 0, 0, 1}) + require.NoError(t, err) + + require.NotPanics(t, func() { + _, err = readPrimitiveArray(Type{ + Type: typeUint8, + IsArray: true, + ArraySize: &arraySize, + }, reader) + }) + require.Error(t, err) +} + +func TestReadMessageRejectsTruncatedSequenceLengthWithoutPanic(t *testing.T) { + t.Parallel() + + msgDefs := map[string]MessageSpecification{ + "test/Message": { + PkgName: "test", + MsgName: "Message", + Fields: []Field{{ + Type: Type{Type: typeUint8, IsArray: true}, + Name: "values", + }}, + }, + } + data := []byte{0, 1, 0, 0, 2, 0} + + var err error + require.NotPanics(t, func() { + _, err = ReadMessage("test/Message", msgDefs, data) + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "past end of data") +} + +func TestScalarReadersReturnErrorsForTruncatedData(t *testing.T) { + t.Parallel() + + tests := map[string]func(*CdrReader) error{ + "boolean": func(reader *CdrReader) error { + _, err := reader.Boolean() + return err + }, + "int8": func(reader *CdrReader) error { + _, err := reader.Int8() + return err + }, + "uint8": func(reader *CdrReader) error { + _, err := reader.Uint8() + return err + }, + "int16": func(reader *CdrReader) error { + _, err := reader.Int16() + return err + }, + "uint16": func(reader *CdrReader) error { + _, err := reader.Uint16() + return err + }, + "int32": func(reader *CdrReader) error { + _, err := reader.Int32() + return err + }, + "uint32": func(reader *CdrReader) error { + _, err := reader.Uint32() + return err + }, + "int64": func(reader *CdrReader) error { + _, err := reader.Int64() + return err + }, + "uint64": func(reader *CdrReader) error { + _, err := reader.Uint64() + return err + }, + "float32": func(reader *CdrReader) error { + _, err := reader.Float32() + return err + }, + "float64": func(reader *CdrReader) error { + _, err := reader.Float64() + return err + }, + } + + for name, read := range tests { + t.Run(name, func(t *testing.T) { + t.Parallel() + + reader, err := NewCdrReader([]byte{0, 1, 0, 0}) + require.NoError(t, err) + + require.NotPanics(t, func() { + err = read(reader) + }) + require.Error(t, err) + }) + } +} + +func TestAlignedScalarReadReturnsErrorWithoutPanic(t *testing.T) { + t.Parallel() + + reader, err := NewCdrReader([]byte{0, 1, 0, 0, 1, 2}) + require.NoError(t, err) + _, err = reader.Uint8() + require.NoError(t, err) + + require.NotPanics(t, func() { + _, err = reader.Uint32() + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "offset: 8") +}