diff --git a/cmd/launcher/launcher.go b/cmd/launcher/launcher.go index 9ab1267171..7a1afe3aaf 100644 --- a/cmd/launcher/launcher.go +++ b/cmd/launcher/launcher.go @@ -382,7 +382,7 @@ func runLauncher(ctx context.Context, cancel func(), multiSlogger, systemMultiSl osqueryruntime.WithAugeasLensFunction(augeas.InstallLenses), ) runGroup.Add("osqueryRunner", osqueryRunner.Run, osqueryRunner.Interrupt) - k.SetInstanceQuerier(osqueryRunner) + k.SetInstanceRunner(osqueryRunner) versionInfo := version.Version() k.SystemSlogger().Log(ctx, slog.LevelInfo, diff --git a/ee/agent/knapsack/knapsack.go b/ee/agent/knapsack/knapsack.go index f52c63ce15..e28dad5b17 100644 --- a/ee/agent/knapsack/knapsack.go +++ b/ee/agent/knapsack/knapsack.go @@ -39,7 +39,7 @@ type knapsack struct { slogger, systemSlogger *multislogger.MultiSlogger - querier types.InstanceQuerier + osqRunner types.OsqRunner // This struct is a work in progress, and will be iteratively added to as needs arise. } @@ -87,9 +87,9 @@ func (k *knapsack) AddSlogHandler(handler ...slog.Handler) { k.systemSlogger.AddHandler(handler...) } -// Osquery instance querier -func (k *knapsack) SetInstanceQuerier(q types.InstanceQuerier) { - k.querier = q +// Osquery instance runner +func (k *knapsack) SetInstanceRunner(r types.OsqRunner) { + k.osqRunner = r } // RegistrationTracker interface methods @@ -97,13 +97,21 @@ func (k *knapsack) RegistrationIDs() []string { return []string{types.DefaultRegistrationID} } +func (k *knapsack) SetRegistrationIDs(registrationIDs []string) error { + if k.osqRunner == nil { + return nil + } + + return k.osqRunner.UpdateRegistrationIDs(registrationIDs) +} + // InstanceStatuses returns the current status of each osquery instance. // It performs a healthcheck against each existing instance. func (k *knapsack) InstanceStatuses() map[string]types.InstanceStatus { - if k.querier == nil { + if k.osqRunner == nil { return nil } - return k.querier.InstanceStatuses() + return k.osqRunner.InstanceStatuses() } // BboltDB interface methods diff --git a/ee/agent/types/knapsack.go b/ee/agent/types/knapsack.go index 45a441cb17..d00f7c602b 100644 --- a/ee/agent/types/knapsack.go +++ b/ee/agent/types/knapsack.go @@ -11,7 +11,7 @@ type Knapsack interface { Slogger RegistrationTracker InstanceQuerier - SetInstanceQuerier(q InstanceQuerier) + SetInstanceRunner(r OsqRunner) // LatestOsquerydPath finds the path to the latest osqueryd binary, after accounting for updates. LatestOsquerydPath(ctx context.Context) string // ReadEnrollSecret returns the enroll secret value, checking in various locations. diff --git a/ee/agent/types/mocks/knapsack.go b/ee/agent/types/mocks/knapsack.go index 58bea406df..bfa7688dd6 100644 --- a/ee/agent/types/mocks/knapsack.go +++ b/ee/agent/types/mocks/knapsack.go @@ -1,4 +1,4 @@ -// Code generated by mockery v2.45.0. DO NOT EDIT. +// Code generated by mockery v2.46.0. DO NOT EDIT. package mocks @@ -1570,9 +1570,9 @@ func (_m *Knapsack) SetInsecureTransportTLS(insecure bool) error { return r0 } -// SetInstanceQuerier provides a mock function with given fields: q -func (_m *Knapsack) SetInstanceQuerier(q types.InstanceQuerier) { - _m.Called(q) +// SetInstanceRunner provides a mock function with given fields: r +func (_m *Knapsack) SetInstanceRunner(r types.OsqRunner) { + _m.Called(r) } // SetKolideServerURL provides a mock function with given fields: url @@ -1760,6 +1760,24 @@ func (_m *Knapsack) SetPinnedOsquerydVersion(version string) error { return r0 } +// SetRegistrationIDs provides a mock function with given fields: registrationIDs +func (_m *Knapsack) SetRegistrationIDs(registrationIDs []string) error { + ret := _m.Called(registrationIDs) + + if len(ret) == 0 { + panic("no return value specified for SetRegistrationIDs") + } + + var r0 error + if rf, ok := ret.Get(0).(func([]string) error); ok { + r0 = rf(registrationIDs) + } else { + r0 = ret.Error(0) + } + + return r0 +} + // SetSystrayRestartEnabled provides a mock function with given fields: enabled func (_m *Knapsack) SetSystrayRestartEnabled(enabled bool) error { ret := _m.Called(enabled) diff --git a/ee/agent/types/registration.go b/ee/agent/types/registration.go index 979889e6ce..1c39ef1f78 100644 --- a/ee/agent/types/registration.go +++ b/ee/agent/types/registration.go @@ -9,4 +9,5 @@ const ( // data may be provided by e.g. a control server subsystem. type RegistrationTracker interface { RegistrationIDs() []string + SetRegistrationIDs(registrationIDs []string) error } diff --git a/ee/agent/types/runner.go b/ee/agent/types/runner.go new file mode 100644 index 0000000000..36768fba6d --- /dev/null +++ b/ee/agent/types/runner.go @@ -0,0 +1,13 @@ +package types + +type ( + // RegistrationChangeHandler is implemented by pkg/osquery/runtime/runner.go + RegistrationChangeHandler interface { + UpdateRegistrationIDs(registrationIDs []string) error + } + + OsqRunner interface { + RegistrationChangeHandler + InstanceQuerier + } +) diff --git a/pkg/osquery/runtime/runner.go b/pkg/osquery/runtime/runner.go index 26389955ce..2dd063a705 100644 --- a/pkg/osquery/runtime/runner.go +++ b/pkg/osquery/runtime/runner.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "log/slog" + "slices" "sync" "sync/atomic" "time" @@ -27,6 +28,7 @@ type settingsStoreWriter interface { type Runner struct { registrationIds []string // we expect to run one instance per registration ID + regIDLock sync.Mutex // locks access to registrationIds instances map[string]*OsqueryInstance // maps registration ID to currently-running instance instanceLock sync.Mutex // locks access to `instances` to avoid e.g. restarting an instance that isn't running yet slogger *slog.Logger @@ -34,7 +36,8 @@ type Runner struct { serviceClient service.KolideService // shared service client for communication between osquery instance and Kolide SaaS settingsWriter settingsStoreWriter // writes to startup settings store opts []OsqueryInstanceOption // global options applying to all osquery instances - shutdown chan struct{} + shutdown chan struct{} // buffered shutdown channel to enable shutting down to restart or exit + rerunRequired atomic.Bool interrupted atomic.Bool } @@ -46,8 +49,9 @@ func New(k types.Knapsack, serviceClient service.KolideService, settingsWriter s knapsack: k, serviceClient: serviceClient, settingsWriter: settingsWriter, - shutdown: make(chan struct{}), - opts: opts, + // the buffer length is arbitrarily set at 100, this number just needs to be higher than the total possible instances + shutdown: make(chan struct{}, 100), + opts: opts, } k.RegisterChangeObserver(runner, @@ -57,12 +61,56 @@ func New(k types.Knapsack, serviceClient service.KolideService, settingsWriter s return runner } +// String method is only added to runner because it is often used in our runtime tests as an argument +// passed to mocked knapsack calls. when we AssertExpectations, the runner struct is traversed by the +// Diff logic inside testify. This causes data races to be incorrectly reported for structs containing mutexes- +// the second read is coming from testify. +// see (one of) the issues here for additional context https://github.com/stretchr/testify/issues/1597 +// If we really needed to expose more here, we could acquire all locks and return fmt.Sprintf("%#v", r). but given +// that we do not, it seems safer to avoid introducing any additional lock contention in case our Stringer call +// is invoked someday in a production flow +func (r *Runner) String() string { + return "runtime.Runner{}" +} + func (r *Runner) Run() error { + for { + err := r.runRegisteredInstances() + if err != nil { + // log any errors but continue, in case we intend to reload + r.slogger.Log(context.TODO(), slog.LevelWarn, + "runRegisteredInstances terminated with error", + "err", err, + ) + } + + // if we're in a state that required re-running all registered instances, + // reset the field and do that + if r.rerunRequired.Load() { + r.rerunRequired.Store(false) + continue + } + + return err + } +} + +func (r *Runner) runRegisteredInstances() error { + // clear the internal instances to add back in fresh as we runInstance, + // this prevents old instances from sticking around if a registrationID is ever removed + r.instanceLock.Lock() + r.instances = make(map[string]*OsqueryInstance) + r.instanceLock.Unlock() + // Create a group to track the workers running each instance wg, ctx := errgroup.WithContext(context.TODO()) // Start each worker for each instance - for _, registrationId := range r.registrationIds { + r.regIDLock.Lock() + regIDs := r.registrationIds + r.regIDLock.Unlock() + + for _, registrationId := range regIDs { id := registrationId wg.Go(func() error { if err := r.runInstance(id); err != nil { @@ -205,6 +253,13 @@ func (r *Runner) Query(query string) ([]map[string]string, error) { } func (r *Runner) Interrupt(_ error) { + if r.interrupted.Load() { + // Already shut down, nothing else to do + return + } + + r.interrupted.Store(true) + if err := r.Shutdown(); err != nil { r.slogger.Log(context.TODO(), slog.LevelWarn, "could not shut down runner on interrupt", @@ -218,14 +273,12 @@ func (r *Runner) Interrupt(_ error) { func (r *Runner) Shutdown() error { ctx, span := traces.StartSpan(context.TODO()) defer span.End() - - if r.interrupted.Load() { - // Already shut down, nothing else to do - return nil + // ensure one shutdown is sent for each instance to read + r.instanceLock.Lock() + for range r.instances { + r.shutdown <- struct{}{} } - - r.interrupted.Store(true) - close(r.shutdown) + r.instanceLock.Unlock() if err := r.triggerShutdownForInstances(ctx); err != nil { return fmt.Errorf("triggering shutdown for instances during runner shutdown: %w", err) @@ -326,7 +379,11 @@ func (r *Runner) Healthy() error { defer r.instanceLock.Unlock() healthcheckErrs := make([]error, 0) - for _, registrationId := range r.registrationIds { + r.regIDLock.Lock() + regIDs := r.registrationIds + r.regIDLock.Unlock() + + for _, registrationId := range regIDs { instance, ok := r.instances[registrationId] if !ok { healthcheckErrs = append(healthcheckErrs, fmt.Errorf("running instance does not exist for %s", registrationId)) @@ -349,8 +406,11 @@ func (r *Runner) InstanceStatuses() map[string]types.InstanceStatus { r.instanceLock.Lock() defer r.instanceLock.Unlock() + r.regIDLock.Lock() + regIDs := r.registrationIds + r.regIDLock.Unlock() instanceStatuses := make(map[string]types.InstanceStatus) - for _, registrationId := range r.registrationIds { + for _, registrationId := range regIDs { instance, ok := r.instances[registrationId] if !ok { instanceStatuses[registrationId] = types.InstanceStatusNotStarted @@ -367,3 +427,47 @@ func (r *Runner) InstanceStatuses() map[string]types.InstanceStatus { return instanceStatuses } + +// UpdateRegistrationIDs detects any changes between the new and stored registration IDs, +// and resets the runner instances for the new registrationIDs if required +func (r *Runner) UpdateRegistrationIDs(newRegistrationIDs []string) error { + slices.Sort(newRegistrationIDs) + + r.regIDLock.Lock() + existingRegistrationIDs := r.registrationIds + r.regIDLock.Unlock() + slices.Sort(existingRegistrationIDs) + + if slices.Equal(newRegistrationIDs, existingRegistrationIDs) { + r.slogger.Log(context.TODO(), slog.LevelDebug, + "skipping runner restarts for updated registration IDs, no changes detected", + ) + + return nil + } + + r.slogger.Log(context.TODO(), slog.LevelDebug, + "detected changes to registrationIDs, will restart runner instances", + "previous_registration_ids", existingRegistrationIDs, + "new_registration_ids", newRegistrationIDs, + ) + + // we know there are changes, safe to update the internal registrationIDs now + r.regIDLock.Lock() + r.registrationIds = newRegistrationIDs + r.regIDLock.Unlock() + // mark rerun as required so that we can safely shutdown all workers and have the changes + // picked back up from within the main Run function + r.rerunRequired.Store(true) + + if err := r.Shutdown(); err != nil { + r.slogger.Log(context.TODO(), slog.LevelWarn, + "could not shut down runner instances for restart after registration changes", + "err", err, + ) + + return err + } + + return nil +} diff --git a/pkg/osquery/runtime/runtime_test.go b/pkg/osquery/runtime/runtime_test.go index 8d89df8193..374d6dd947 100644 --- a/pkg/osquery/runtime/runtime_test.go +++ b/pkg/osquery/runtime/runtime_test.go @@ -210,7 +210,6 @@ func TestWithOsqueryFlags(t *testing.T) { func TestFlagsChanged(t *testing.T) { t.Parallel() - rootDirectory := testRootDirectory(t) logBytes, slogger := setUpTestSlogger() @@ -781,6 +780,180 @@ func TestRestart(t *testing.T) { waitShutdown(t, runner, logBytes) } +func TestMultipleInstancesWithUpdatedRegistrationIDs(t *testing.T) { + t.Parallel() + rootDirectory := testRootDirectory(t) + + logBytes, slogger := setUpTestSlogger() + + k := typesMocks.NewKnapsack(t) + k.On("RegistrationIDs").Return([]string{types.DefaultRegistrationID}) + k.On("OsqueryHealthcheckStartupDelay").Return(0 * time.Second).Maybe() + k.On("WatchdogEnabled").Return(false) + k.On("RegisterChangeObserver", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything) + k.On("Slogger").Return(slogger) + k.On("LatestOsquerydPath", mock.Anything).Return(testOsqueryBinary) + k.On("RootDirectory").Return(rootDirectory).Maybe() + k.On("OsqueryFlags").Return([]string{}) + k.On("OsqueryVerbose").Return(true) + k.On("LoggingInterval").Return(5 * time.Minute).Maybe() + k.On("LogMaxBytesPerBatch").Return(0).Maybe() + k.On("Transport").Return("jsonrpc").Maybe() + k.On("ReadEnrollSecret").Return("", nil).Maybe() + k.On("InModernStandby").Return(false).Maybe() + k.On("RegisterChangeObserver", mock.Anything, keys.UpdateChannel).Maybe() + k.On("RegisterChangeObserver", mock.Anything, keys.PinnedLauncherVersion).Maybe() + k.On("RegisterChangeObserver", mock.Anything, keys.PinnedOsquerydVersion).Maybe() + k.On("UpdateChannel").Return("stable").Maybe() + k.On("PinnedLauncherVersion").Return("").Maybe() + k.On("PinnedOsquerydVersion").Return("").Maybe() + setUpMockStores(t, k) + serviceClient := mockServiceClient(t) + + storeWriter := settingsstoremock.NewSettingsStoreWriter(t) + storeWriter.On("WriteSettings").Return(nil) + runner := New(k, serviceClient, storeWriter) + ensureShutdownOnCleanup(t, runner, logBytes) + + // Start the instance + go runner.Run() + waitHealthy(t, runner, logBytes) + + // Confirm the default instance was started + require.Contains(t, runner.instances, types.DefaultRegistrationID) + require.NotNil(t, runner.instances[types.DefaultRegistrationID].stats) + require.NotEmpty(t, runner.instances[types.DefaultRegistrationID].stats.StartTime, "start time should be added to default instance stats on start up") + require.NotEmpty(t, runner.instances[types.DefaultRegistrationID].stats.ConnectTime, "connect time should be added to default instance stats on start up") + + // confirm only the default instance has started + require.Equal(t, 1, len(runner.instances)) + + // Confirm instance statuses are reported correctly + instanceStatuses := runner.InstanceStatuses() + require.Contains(t, instanceStatuses, types.DefaultRegistrationID) + require.Equal(t, instanceStatuses[types.DefaultRegistrationID], types.InstanceStatusHealthy) + + // Add in an extra instance + extraRegistrationId := ulid.New() + updateErr := runner.UpdateRegistrationIDs([]string{types.DefaultRegistrationID, extraRegistrationId}) + require.NoError(t, updateErr) + waitHealthy(t, runner, logBytes) + updatedInstanceStatuses := runner.InstanceStatuses() + // verify that rerunRequired has been reset for any future changes + require.False(t, runner.rerunRequired.Load()) + // now verify both instances are reported + require.Equal(t, 2, len(runner.instances)) + require.Contains(t, updatedInstanceStatuses, types.DefaultRegistrationID) + require.Contains(t, updatedInstanceStatuses, extraRegistrationId) + // Confirm the additional instance was started and is healthy + require.NotNil(t, runner.instances[extraRegistrationId].stats) + require.NotEmpty(t, runner.instances[extraRegistrationId].stats.StartTime, "start time should be added to secondary instance stats on start up") + require.NotEmpty(t, runner.instances[extraRegistrationId].stats.ConnectTime, "connect time should be added to secondary instance stats on start up") + require.Equal(t, updatedInstanceStatuses[extraRegistrationId], types.InstanceStatusHealthy) + + // update registration IDs one more time, this time removing the additional registration + originalDefaultInstanceStartTime := runner.instances[extraRegistrationId].stats.StartTime + updateErr = runner.UpdateRegistrationIDs([]string{types.DefaultRegistrationID}) + require.NoError(t, updateErr) + waitHealthy(t, runner, logBytes) + + // now verify only the default instance remains + require.Equal(t, 1, len(runner.instances)) + // Confirm the default instance was started and is healthy + require.Contains(t, runner.instances, types.DefaultRegistrationID) + require.NotNil(t, runner.instances[types.DefaultRegistrationID].stats) + require.NotEmpty(t, runner.instances[types.DefaultRegistrationID].stats.StartTime, "start time should be added to default instance stats on start up") + require.NotEmpty(t, runner.instances[types.DefaultRegistrationID].stats.ConnectTime, "connect time should be added to default instance stats on start up") + // verify that rerunRequired has been reset for any future changes + require.False(t, runner.rerunRequired.Load()) + // verify the default instance was restarted + require.NotEqual(t, originalDefaultInstanceStartTime, runner.instances[types.DefaultRegistrationID].stats.StartTime) + + waitShutdown(t, runner, logBytes) + + // Confirm instance exited + require.Contains(t, runner.instances, types.DefaultRegistrationID) + require.NotNil(t, runner.instances[types.DefaultRegistrationID].stats) + require.NotEmpty(t, runner.instances[types.DefaultRegistrationID].stats.ExitTime, "exit time should be added to default instance stats on shutdown") +} + +func TestUpdatingRegistrationIDsOnlyRestartsForChanges(t *testing.T) { + t.Parallel() + rootDirectory := testRootDirectory(t) + + logBytes, slogger := setUpTestSlogger() + extraRegistrationId := ulid.New() + + k := typesMocks.NewKnapsack(t) + k.On("RegistrationIDs").Return([]string{types.DefaultRegistrationID, extraRegistrationId}) + k.On("OsqueryHealthcheckStartupDelay").Return(0 * time.Second).Maybe() + k.On("WatchdogEnabled").Return(false) + k.On("RegisterChangeObserver", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything) + k.On("Slogger").Return(slogger) + k.On("LatestOsquerydPath", mock.Anything).Return(testOsqueryBinary) + k.On("RootDirectory").Return(rootDirectory).Maybe() + k.On("OsqueryFlags").Return([]string{}) + k.On("OsqueryVerbose").Return(true) + k.On("LoggingInterval").Return(5 * time.Minute).Maybe() + k.On("LogMaxBytesPerBatch").Return(0).Maybe() + k.On("Transport").Return("jsonrpc").Maybe() + k.On("ReadEnrollSecret").Return("", nil).Maybe() + k.On("InModernStandby").Return(false).Maybe() + k.On("RegisterChangeObserver", mock.Anything, keys.UpdateChannel).Maybe() + k.On("RegisterChangeObserver", mock.Anything, keys.PinnedLauncherVersion).Maybe() + k.On("RegisterChangeObserver", mock.Anything, keys.PinnedOsquerydVersion).Maybe() + k.On("UpdateChannel").Return("stable").Maybe() + k.On("PinnedLauncherVersion").Return("").Maybe() + k.On("PinnedOsquerydVersion").Return("").Maybe() + setUpMockStores(t, k) + serviceClient := mockServiceClient(t) + + storeWriter := settingsstoremock.NewSettingsStoreWriter(t) + storeWriter.On("WriteSettings").Return(nil) + runner := New(k, serviceClient, storeWriter) + ensureShutdownOnCleanup(t, runner, logBytes) + + // Start the instance + go runner.Run() + waitHealthy(t, runner, logBytes) + + require.Equal(t, 2, len(runner.instances)) + // Confirm the default instance was started + require.Contains(t, runner.instances, types.DefaultRegistrationID) + require.NotNil(t, runner.instances[types.DefaultRegistrationID].stats) + require.NotEmpty(t, runner.instances[types.DefaultRegistrationID].stats.StartTime, "start time should be added to default instance stats on start up") + require.NotEmpty(t, runner.instances[types.DefaultRegistrationID].stats.ConnectTime, "connect time should be added to default instance stats on start up") + // note the original start time + defaultInstanceStartTime := runner.instances[types.DefaultRegistrationID].stats.StartTime + + // Confirm the extra instance was started + require.Contains(t, runner.instances, extraRegistrationId) + require.NotNil(t, runner.instances[extraRegistrationId].stats) + require.NotEmpty(t, runner.instances[extraRegistrationId].stats.StartTime, "start time should be added to extra instance stats on start up") + require.NotEmpty(t, runner.instances[extraRegistrationId].stats.ConnectTime, "connect time should be added to extra instance stats on start up") + // note the original start time + extraInstanceStartTime := runner.instances[extraRegistrationId].stats.StartTime + + // rerun with identical registrationIDs in swapped order and verify that the instances are not restarted + updateErr := runner.UpdateRegistrationIDs([]string{extraRegistrationId, types.DefaultRegistrationID}) + require.NoError(t, updateErr) + waitHealthy(t, runner, logBytes) + + require.Equal(t, 2, len(runner.instances)) + require.Equal(t, extraInstanceStartTime, runner.instances[extraRegistrationId].stats.StartTime) + require.Equal(t, defaultInstanceStartTime, runner.instances[types.DefaultRegistrationID].stats.StartTime) + + waitShutdown(t, runner, logBytes) + + // Confirm both instances exited + require.Contains(t, runner.instances, types.DefaultRegistrationID) + require.NotNil(t, runner.instances[types.DefaultRegistrationID].stats) + require.NotEmpty(t, runner.instances[types.DefaultRegistrationID].stats.ExitTime, "exit time should be added to default instance stats on shutdown") + require.Contains(t, runner.instances, extraRegistrationId) + require.NotNil(t, runner.instances[extraRegistrationId].stats) + require.NotEmpty(t, runner.instances[extraRegistrationId].stats.ExitTime, "exit time should be added to secondary instance stats on shutdown") +} + // sets up an osquery instance with a running extension to be used in tests. func setupOsqueryInstanceForTests(t *testing.T) (runner *Runner, logBytes *threadsafebuffer.ThreadSafeBuffer) { rootDirectory := testRootDirectory(t)