Distributed lock and store calendar_events UUID as binary in MySQL (#20277)

#19352

Fix for code review comment:
https://github.com/fleetdm/fleet/pull/20156#discussion_r1668421504

Also includes changes from https://github.com/fleetdm/fleet/pull/20252

# Checklist for submitter

If some of the following don't apply, delete the relevant line.

<!-- Note that API documentation changes are now addressed by the
product design team. -->

- [x] Added/updated tests
- [x] If database migrations are included, checked table schema to
confirm autoupdate
- For database migrations:
- [x] Checked schema for all modified table for columns that will
auto-update timestamps during migration.
- [x] Confirmed that updating the timestamps is acceptable, and will not
cause unwanted side effects.
- [x] Ensured the correct collation is explicitly set for character
columns (`COLLATE utf8mb4_unicode_ci`).
- [x] Manual QA for all new/changed functionality
This commit is contained in:
Victor Lyuboslavsky
2024-07-10 08:49:05 -05:00
committed by GitHub
parent 886ab9098d
commit 7bcd61a8bd
21 changed files with 964 additions and 79 deletions
+5 -1
View File
@@ -50,6 +50,7 @@ import (
"github.com/fleetdm/fleet/v4/server/pubsub"
"github.com/fleetdm/fleet/v4/server/service"
"github.com/fleetdm/fleet/v4/server/service/async"
"github.com/fleetdm/fleet/v4/server/service/redis_lock"
"github.com/fleetdm/fleet/v4/server/service/redis_policy_set"
"github.com/fleetdm/fleet/v4/server/sso"
"github.com/fleetdm/fleet/v4/server/version"
@@ -691,6 +692,7 @@ the way that the Fleet server works.
}
var softwareInstallStore fleet.SoftwareInstallerStore
var distributedLock fleet.Lock
if license.IsPremium() {
profileMatcher := apple_mdm.NewProfileMatcher(redisPool)
if config.S3.SoftwareInstallersBucket != "" {
@@ -718,6 +720,7 @@ the way that the Fleet server works.
}
}
distributedLock = redis_lock.NewLock(redisPool)
svc, err = eeservice.NewService(
svc,
ds,
@@ -730,6 +733,7 @@ the way that the Fleet server works.
ssoSessionStore,
profileMatcher,
softwareInstallStore,
distributedLock,
)
if err != nil {
initFatal(err, "initial Fleet Premium service")
@@ -870,7 +874,7 @@ the way that the Fleet server works.
} else {
config.Calendar.Periodicity = 5 * time.Minute
}
return cron.NewCalendarSchedule(ctx, instanceID, ds, config.Calendar, logger)
return cron.NewCalendarSchedule(ctx, instanceID, ds, distributedLock, config.Calendar, logger)
},
); err != nil {
initFatal(err, "failed to register calendar schedule")
+214 -6
View File
@@ -3,16 +3,24 @@ package service
import (
"context"
"fmt"
"sync"
"github.com/fleetdm/fleet/v4/server/authz"
"github.com/fleetdm/fleet/v4/server/contexts/ctxerr"
"github.com/fleetdm/fleet/v4/server/fleet"
"github.com/fleetdm/fleet/v4/server/service/calendar"
"github.com/go-kit/log/level"
"sync"
"github.com/google/uuid"
)
var asyncCalendarProcessing bool
var asyncMutex sync.Mutex
func (svc *Service) CalendarWebhook(ctx context.Context, eventUUID string, channelID string, resourceState string) error {
// We don't want the sender to cancel the context since we want to make sure we process the webhook.
ctx = context.WithoutCancel(ctx)
appConfig, err := svc.ds.AppConfig(ctx)
if err != nil {
return fmt.Errorf("load app config: %w", err)
@@ -36,9 +44,9 @@ func (svc *Service) CalendarWebhook(ctx context.Context, eventUUID string, chann
svc.authz.SkipAuthorization(ctx)
if fleet.IsNotFound(err) {
// We could try to stop the channel callbacks here, but that may not be secure since we don't know if the request is legitimate
level.Warn(svc.logger).Log("msg", "Received calendar callback, but did not find corresponding event in database", "event_uuid",
level.Info(svc.logger).Log("msg", "Received calendar callback, but did not find corresponding event in database", "event_uuid",
eventUUID, "channel_id", channelID)
return err
return nil
}
return err
}
@@ -47,7 +55,7 @@ func (svc *Service) CalendarWebhook(ctx context.Context, eventUUID string, chann
return fmt.Errorf("calendar event %s has no team ID", eventUUID)
}
localConfig := &calendar.CalendarConfig{
localConfig := &calendar.Config{
GoogleCalendarIntegration: *googleCalendarIntegrationConfig,
ServerURL: appConfig.ServerSettings.ServerURL,
}
@@ -63,6 +71,60 @@ func (svc *Service) CalendarWebhook(ctx context.Context, eventUUID string, chann
return authz.ForbiddenWithInternal(fmt.Sprintf("calendar channel ID mismatch: %s != %s", savedChannelID, channelID), nil, nil, nil)
}
lockValue, reserved, err := svc.getCalendarLock(ctx, eventUUID, true)
if err != nil {
return err
}
// If lock has been reserved by cron, we will need to re-process this event in case the calendar event was changed after the cron job read it.
if lockValue == "" && !reserved {
// We did not get a lock, so there is nothing to do here
return nil
}
if !reserved {
unlocked := false
defer func() {
if !unlocked {
svc.releaseCalendarLock(ctx, eventUUID, lockValue)
}
}()
// Remove event from the queue so that we don't process this event again.
// Note: This item can be added back to the queue while we are processing it.
err = svc.distributedLock.RemoveFromSet(ctx, calendar.QueueKey, eventUUID)
if err != nil {
return ctxerr.Wrap(ctx, err, "remove calendar event from queue")
}
err = svc.processCalendarEvent(ctx, eventDetails, googleCalendarIntegrationConfig, userCalendar)
if err != nil {
return err
}
svc.releaseCalendarLock(ctx, eventUUID, lockValue)
unlocked = true
}
// Now, we need to check if there are any events in the queue that need to be re-processed.
asyncMutex.Lock()
defer asyncMutex.Unlock()
if !asyncCalendarProcessing {
eventIDs, err := svc.distributedLock.GetSet(ctx, calendar.QueueKey)
if err != nil {
return ctxerr.Wrap(ctx, err, "get calendar event queue")
}
if len(eventIDs) > 0 {
asyncCalendarProcessing = true
go svc.processCalendarAsync(ctx, eventIDs)
}
return nil
}
return nil
}
func (svc *Service) processCalendarEvent(ctx context.Context, eventDetails *fleet.CalendarEventDetails,
googleCalendarIntegrationConfig *fleet.GoogleCalendarIntegration, userCalendar fleet.UserCalendar) error {
genBodyFn := func(conflict bool) (body string, ok bool, err error) {
// This function is called when a new event is being created.
@@ -113,7 +175,7 @@ func (svc *Service) CalendarWebhook(ctx context.Context, eventUUID string, chann
return calendar.GenerateCalendarEventBody(ctx, svc.ds, team.Name, host, &sync.Map{}, conflict, svc.logger), true, nil
}
err = userCalendar.Configure(eventDetails.Email)
err := userCalendar.Configure(eventDetails.Email)
if err != nil {
return ctxerr.Wrap(ctx, err, "configure calendar")
}
@@ -124,10 +186,156 @@ func (svc *Service) CalendarWebhook(ctx context.Context, eventUUID string, chann
if updated && event != nil {
// Event was updated, so we need to save it
_, err = svc.ds.CreateOrUpdateCalendarEvent(ctx, event.UUID, event.Email, event.StartTime, event.EndTime, event.Data,
event.TimeZone, eventDetails.ID, fleet.CalendarWebhookStatusNone)
event.TimeZone, eventDetails.HostID, fleet.CalendarWebhookStatusNone)
if err != nil {
return ctxerr.Wrap(ctx, err, "create or update calendar event")
}
}
return nil
}
func (svc *Service) releaseCalendarLock(ctx context.Context, eventUUID string, lockValue string) {
ok, err := svc.distributedLock.ReleaseLock(ctx, calendar.LockKeyPrefix+eventUUID, lockValue)
if err != nil {
level.Error(svc.logger).Log("msg", "Failed to release calendar lock", "err", err)
}
if !ok {
// If the lock was not released, it will expire on its own.
level.Warn(svc.logger).Log("msg", "Failed to release calendar lock")
}
}
func (svc *Service) getCalendarLock(ctx context.Context, eventUUID string, addToQueue bool) (lockValue string, reserved bool, err error) {
// Check if lock has been reserved, which means we can't have it.
reservedValue, err := svc.distributedLock.Get(ctx, calendar.ReservedLockKeyPrefix+eventUUID)
if err != nil {
return "", false, ctxerr.Wrap(ctx, err, "get calendar reserved lock")
}
reserved = reservedValue != nil
if reserved && !addToQueue {
// We flag the lock as reserved.
return "", reserved, nil
}
var lockAcquired bool
if !reserved {
// Try to acquire the lock
lockValue = uuid.New().String()
lockAcquired, err = svc.distributedLock.AcquireLock(ctx, calendar.LockKeyPrefix+eventUUID, lockValue, 0)
if err != nil {
return "", false, ctxerr.Wrap(ctx, err, "acquire calendar lock")
}
}
if (!lockAcquired || reserved) && addToQueue {
// Could not acquire lock, so we are already processing this event. In this case, we add the event to
// the queue (actually a set) to indicate that we need to re-process the event.
err = svc.distributedLock.AddToSet(ctx, calendar.QueueKey, eventUUID)
if err != nil {
return "", false, ctxerr.Wrap(ctx, err, "add calendar event to queue")
}
if reserved {
// We flag the lock as reserved.
return "", reserved, nil
}
// Try to acquire the lock again in case it was released while we were adding the event to the queue.
lockAcquired, err = svc.distributedLock.AcquireLock(ctx, calendar.LockKeyPrefix+eventUUID, lockValue, 0)
if err != nil {
return "", false, ctxerr.Wrap(ctx, err, "acquire calendar lock again")
}
if !lockAcquired {
// We could not acquire the lock, so we are done here.
return "", reserved, nil
}
}
return lockValue, false, nil
}
func (svc *Service) processCalendarAsync(ctx context.Context, eventIDs []string) {
defer func() {
asyncMutex.Lock()
asyncCalendarProcessing = false
asyncMutex.Unlock()
}()
for {
if len(eventIDs) == 0 {
return
}
for _, eventUUID := range eventIDs {
if ok := svc.processCalendarEventAsync(ctx, eventUUID); !ok {
return
}
}
// Now we check whether there are any more events in the queue.
var err error
eventIDs, err = svc.distributedLock.GetSet(ctx, calendar.QueueKey)
if err != nil {
level.Error(svc.logger).Log("msg", "Failed to get calendar event queue", "err", err)
return
}
}
}
func (svc *Service) processCalendarEventAsync(ctx context.Context, eventUUID string) bool {
lockValue, _, err := svc.getCalendarLock(ctx, eventUUID, false)
if err != nil {
level.Error(svc.logger).Log("msg", "Failed to get calendar lock", "err", err)
return false
}
if lockValue == "" {
// We did not get a lock, so there is nothing to do here
return true
}
defer svc.releaseCalendarLock(ctx, eventUUID, lockValue)
// Remove event from the queue so that we don't process this event again.
// Note: This item can be added back to the queue while we are processing it.
err = svc.distributedLock.RemoveFromSet(ctx, calendar.QueueKey, eventUUID)
if err != nil {
level.Error(svc.logger).Log("msg", "Failed to remove calendar event from queue", "err", err)
return false
}
appConfig, err := svc.ds.AppConfig(ctx)
if err != nil {
level.Error(svc.logger).Log("msg", "Failed to load app config", "err", err)
return false
}
if len(appConfig.Integrations.GoogleCalendar) == 0 {
// Google Calendar integration is not configured
return true
}
googleCalendarIntegrationConfig := appConfig.Integrations.GoogleCalendar[0]
eventDetails, err := svc.ds.GetCalendarEventDetailsByUUID(ctx, eventUUID)
if err != nil {
if fleet.IsNotFound(err) {
// We found this event when the callback initially came in. So the event may have been removed or re-created since then.
return true
}
level.Error(svc.logger).Log("msg", "Failed to get calendar event details", "err", err)
return false
}
if eventDetails.TeamID == nil {
// Should not happen
level.Error(svc.logger).Log("msg", "Calendar event has no team ID", "uuid", eventUUID)
return false
}
localConfig := &calendar.Config{
GoogleCalendarIntegration: *googleCalendarIntegrationConfig,
ServerURL: appConfig.ServerSettings.ServerURL,
}
userCalendar := calendar.CreateUserCalendarFromConfig(ctx, localConfig, svc.logger)
err = svc.processCalendarEvent(ctx, eventDetails, googleCalendarIntegrationConfig, userCalendar)
if err != nil {
level.Error(svc.logger).Log("msg", "Failed to process calendar event", "err", err)
return false
}
return true
}
+1
View File
@@ -82,6 +82,7 @@ func setupMockDatastorePremiumService() (*mock.Store, *eeservice.Service, contex
nil,
nil,
nil,
nil,
)
if err != nil {
panic(err)
+3
View File
@@ -28,6 +28,7 @@ type Service struct {
depService *apple_mdm.DEPService
profileMatcher fleet.ProfileMatcher
softwareInstallStore fleet.SoftwareInstallerStore
distributedLock fleet.Lock
}
func NewService(
@@ -42,6 +43,7 @@ func NewService(
sso sso.SessionStore,
profileMatcher fleet.ProfileMatcher,
softwareInstallStore fleet.SoftwareInstallerStore,
distributedLock fleet.Lock,
) (*Service, error) {
authorizer, err := authz.NewAuthorizer()
if err != nil {
@@ -61,6 +63,7 @@ func NewService(
depService: apple_mdm.NewDEPService(ds, depStorage, logger),
profileMatcher: profileMatcher,
softwareInstallStore: softwareInstallStore,
distributedLock: distributedLock,
}
// Override methods that can't be easily overriden via
+95 -15
View File
@@ -4,6 +4,7 @@ import (
"context"
"errors"
"fmt"
"github.com/google/uuid"
"slices"
"sync"
"time"
@@ -27,6 +28,7 @@ func NewCalendarSchedule(
ctx context.Context,
instanceID string,
ds fleet.Datastore,
distributedLock fleet.Lock,
serverConfig config.CalendarConfig,
logger kitlog.Logger,
) (*schedule.Schedule, error) {
@@ -47,7 +49,7 @@ func NewCalendarSchedule(
schedule.WithJob(
"calendar_events",
func(ctx context.Context) error {
return cronCalendarEvents(ctx, ds, serverConfig, logger)
return cronCalendarEvents(ctx, ds, distributedLock, serverConfig, logger)
},
),
)
@@ -55,7 +57,8 @@ func NewCalendarSchedule(
return s, nil
}
func cronCalendarEvents(ctx context.Context, ds fleet.Datastore, serverConfig config.CalendarConfig, logger kitlog.Logger) error {
func cronCalendarEvents(ctx context.Context, ds fleet.Datastore, distributedLock fleet.Lock, serverConfig config.CalendarConfig,
logger kitlog.Logger) error {
appConfig, err := ds.AppConfig(ctx)
if err != nil {
return fmt.Errorf("load app config: %w", err)
@@ -78,14 +81,14 @@ func cronCalendarEvents(ctx context.Context, ds fleet.Datastore, serverConfig co
return fmt.Errorf("list teams: %w", err)
}
localConfig := &calendar.CalendarConfig{
localConfig := &calendar.Config{
CalendarConfig: serverConfig,
GoogleCalendarIntegration: *googleCalendarIntegrationConfig,
ServerURL: appConfig.ServerSettings.ServerURL,
}
for _, team := range teams {
if err := cronCalendarEventsForTeam(
ctx, ds, localConfig, *team, appConfig.OrgInfo.OrgName, domain, logger,
ctx, ds, distributedLock, localConfig, *team, appConfig.OrgInfo.OrgName, domain, logger,
); err != nil {
level.Info(logger).Log("msg", "events calendar cron", "team_id", team.ID, "err", err)
}
@@ -97,7 +100,8 @@ func cronCalendarEvents(ctx context.Context, ds fleet.Datastore, serverConfig co
func cronCalendarEventsForTeam(
ctx context.Context,
ds fleet.Datastore,
calendarConfig *calendar.CalendarConfig,
distributedLock fleet.Lock,
calendarConfig *calendar.Config,
team fleet.Team,
orgName string,
domain string,
@@ -179,7 +183,7 @@ func cronCalendarEventsForTeam(
// Process hosts that are failing calendar policies.
start = time.Now()
processCalendarFailingHosts(ctx, ds, calendarConfig, orgName, failingHosts, logger)
processCalendarFailingHosts(ctx, ds, distributedLock, calendarConfig, orgName, failingHosts, logger)
level.Debug(logger).Log(
"msg", "failing_hosts", "took", time.Since(start),
)
@@ -197,7 +201,8 @@ func cronCalendarEventsForTeam(
func processCalendarFailingHosts(
ctx context.Context,
ds fleet.Datastore,
calendarConfig *calendar.CalendarConfig,
distributedLock fleet.Lock,
calendarConfig *calendar.Config,
orgName string,
hosts []fleet.HostPolicyMembershipData,
logger kitlog.Logger,
@@ -253,7 +258,8 @@ func processCalendarFailingHosts(
switch {
case err == nil && !expiredEvent:
if err := processFailingHostExistingCalendarEvent(
ctx, ds, userCalendar, orgName, hostCalendarEvent, calendarEvent, host, &policyIDtoPolicy, calendarConfig, logger,
ctx, ds, distributedLock, userCalendar, orgName, hostCalendarEvent, calendarEvent, host, &policyIDtoPolicy,
calendarConfig, logger,
); err != nil {
level.Info(logger).Log("msg", "process failing host existing calendar event", "err", err)
continue // continue with next host
@@ -303,15 +309,89 @@ func filterHostsWithSameEmail(hosts []fleet.HostPolicyMembershipData) []fleet.Ho
func processFailingHostExistingCalendarEvent(
ctx context.Context,
ds fleet.Datastore,
distributedLock fleet.Lock,
userCalendar fleet.UserCalendar,
orgName string,
hostCalendarEvent *fleet.HostCalendarEvent,
calendarEvent *fleet.CalendarEvent,
host fleet.HostPolicyMembershipData,
policyIDtoPolicy *sync.Map,
calendarConfig *calendar.CalendarConfig,
calendarConfig *calendar.Config,
logger kitlog.Logger,
) error {
// Try to acquire the lock. Lock is needed to ensure calendar callback is not processed for this event at the same time.
eventUUID := calendarEvent.UUID
lockValue := uuid.New().String()
lockAcquired, err := distributedLock.AcquireLock(ctx, calendar.LockKeyPrefix+eventUUID, lockValue, 0)
if err != nil {
return fmt.Errorf("acquire calendar lock: %w", err)
}
lockReserved := false
if !lockAcquired {
// Lock was not acquired. We reserve the lock and try to acquire it until we do.
var timeoutMs uint64 = 2 * 60 * 1000
lockAcquired, err = distributedLock.AcquireLock(ctx, calendar.ReservedLockKeyPrefix+eventUUID, lockValue, timeoutMs)
if err != nil {
return fmt.Errorf("reserve calendar lock: %w", err)
}
if !lockAcquired {
// Lock was not reserved. Another cron job is processing this event. This is not expected.
return errors.New("could not reserve calendar lock")
}
lockReserved = true
done := make(chan struct{})
go func() {
for {
// Keep trying to get the lock.
lockAcquired, err = distributedLock.AcquireLock(ctx, calendar.LockKeyPrefix+eventUUID, lockValue, 0)
if err != nil || lockAcquired {
done <- struct{}{}
return
}
time.Sleep(100 * time.Millisecond)
}
}()
select {
case <-done:
// Lock was acquired.
if err != nil {
return fmt.Errorf("try to acquire calendar lock: %w", err)
}
case <-time.After(time.Duration(timeoutMs) * time.Millisecond):
// We couldn't acquire the lock in time.
return errors.New("could not acquire calendar lock in time")
}
}
defer func() {
// Release locks.
if lockReserved {
ok, err := distributedLock.ReleaseLock(ctx, calendar.ReservedLockKeyPrefix+eventUUID, lockValue)
if err != nil {
level.Error(logger).Log("msg", "Failed to release calendar reserve lock", "err", err)
}
if !ok {
// If the lock was not released, it will expire on its own.
level.Warn(logger).Log("msg", "Failed to release calendar reserve lock")
}
}
ok, err := distributedLock.ReleaseLock(ctx, calendar.LockKeyPrefix+eventUUID, lockValue)
if err != nil {
level.Error(logger).Log("msg", "Failed to release calendar lock", "err", err)
}
if !ok {
// If the lock was not released, it will expire on its own.
level.Warn(logger).Log("msg", "Failed to release calendar lock")
}
}()
// Remove event from the queue so that we don't process this event again.
// Note: This item can be added back to the queue while we are processing it.
err = distributedLock.RemoveFromSet(ctx, calendar.QueueKey, eventUUID)
if err != nil {
return fmt.Errorf("remove calendar event from queue: %w", err)
}
updatedEvent := calendarEvent
updated := false
now := time.Now()
@@ -492,7 +572,7 @@ func addBusinessDay(date time.Time) time.Time {
func removeCalendarEventsFromPassingHosts(
ctx context.Context,
ds fleet.Datastore,
calendarConfig *calendar.CalendarConfig,
calendarConfig *calendar.Config,
hosts []fleet.HostPolicyMembershipData,
logger kitlog.Logger,
) {
@@ -601,9 +681,9 @@ func cronCalendarEventsCleanup(ctx context.Context, ds fleet.Datastore, logger k
}
var userCalendar fleet.UserCalendar
var calConfig *calendar.CalendarConfig
var calConfig *calendar.Config
if len(appConfig.Integrations.GoogleCalendar) > 0 {
calConfig = &calendar.CalendarConfig{
calConfig = &calendar.Config{
GoogleCalendarIntegration: *appConfig.Integrations.GoogleCalendar[0],
ServerURL: appConfig.ServerSettings.ServerURL,
}
@@ -657,7 +737,7 @@ func cronCalendarEventsCleanup(ctx context.Context, ds fleet.Datastore, logger k
func deleteAllCalendarEvents(
ctx context.Context,
ds fleet.Datastore,
calendarConfig *calendar.CalendarConfig,
calendarConfig *calendar.Config,
teamID *uint,
logger kitlog.Logger,
) error {
@@ -670,7 +750,7 @@ func deleteAllCalendarEvents(
}
func deleteCalendarEventsInParallel(
ctx context.Context, ds fleet.Datastore, calendarConfig *calendar.CalendarConfig, calendarEvents []*fleet.CalendarEvent,
ctx context.Context, ds fleet.Datastore, calendarConfig *calendar.Config, calendarEvents []*fleet.CalendarEvent,
logger kitlog.Logger,
) {
if len(calendarEvents) > 0 {
@@ -703,7 +783,7 @@ func deleteCalendarEventsInParallel(
func cleanupTeamCalendarEvents(
ctx context.Context,
ds fleet.Datastore,
calendarConfig *calendar.CalendarConfig,
calendarConfig *calendar.Config,
team fleet.Team,
logger kitlog.Logger,
) error {
+15 -9
View File
@@ -11,14 +11,15 @@ import (
"testing"
"time"
"github.com/fleetdm/fleet/v4/server/config"
"github.com/fleetdm/fleet/v4/server/ptr"
"github.com/stretchr/testify/assert"
"github.com/fleetdm/fleet/v4/ee/server/calendar"
"github.com/fleetdm/fleet/v4/server/config"
"github.com/fleetdm/fleet/v4/server/datastore/redis/redistest"
"github.com/fleetdm/fleet/v4/server/fleet"
"github.com/fleetdm/fleet/v4/server/mock"
"github.com/fleetdm/fleet/v4/server/ptr"
"github.com/fleetdm/fleet/v4/server/service/redis_lock"
kitlog "github.com/go-kit/log"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
@@ -199,7 +200,8 @@ func TestEventForDifferentHost(t *testing.T) {
return hcEvent, calEvent, nil
}
err := cronCalendarEvents(ctx, ds, defaultCalendarConfig, logger)
pool := redistest.SetupRedis(t, t.Name(), false, false, false)
err := cronCalendarEvents(ctx, ds, redis_lock.NewLock(pool), defaultCalendarConfig, logger)
require.NoError(t, err)
}
@@ -373,7 +375,8 @@ func TestCalendarEventsMultipleHosts(t *testing.T) {
return nil, nil
}
err := cronCalendarEvents(ctx, ds, defaultCalendarConfig, logger)
pool := redistest.SetupRedis(t, t.Name(), false, false, false)
err := cronCalendarEvents(ctx, ds, redis_lock.NewLock(pool), defaultCalendarConfig, logger)
require.NoError(t, err)
eventsMu.Lock()
@@ -662,7 +665,9 @@ func TestCalendarEvents1KHosts(t *testing.T) {
return nil, nil
}
err := cronCalendarEvents(ctx, ds, defaultCalendarConfig, logger)
pool := redistest.SetupRedis(t, t.Name(), false, false, false)
distributedLock := redis_lock.NewLock(pool)
err := cronCalendarEvents(ctx, ds, distributedLock, defaultCalendarConfig, logger)
require.NoError(t, err)
createdCalendarEvents := calendar.ListGoogleMockEvents()
@@ -699,7 +704,7 @@ func TestCalendarEvents1KHosts(t *testing.T) {
return nil
}
err = cronCalendarEvents(ctx, ds, defaultCalendarConfig, logger)
err = cronCalendarEvents(ctx, ds, distributedLock, defaultCalendarConfig, logger)
require.NoError(t, err)
createdCalendarEvents = calendar.ListGoogleMockEvents()
@@ -948,7 +953,8 @@ func TestEventBody(t *testing.T) {
return nil, nil
}
err := cronCalendarEvents(ctx, ds, defaultCalendarConfig, logger)
pool := redistest.SetupRedis(t, t.Name(), false, false, false)
err := cronCalendarEvents(ctx, ds, redis_lock.NewLock(pool), defaultCalendarConfig, logger)
require.NoError(t, err)
numberOfEvents := 7
+33 -22
View File
@@ -5,6 +5,7 @@ import (
"database/sql"
"errors"
"fmt"
"github.com/google/uuid"
"time"
"github.com/fleetdm/fleet/v4/server/contexts/ctxerr"
@@ -12,9 +13,11 @@ import (
"github.com/jmoiron/sqlx"
)
const calendarEventCols = `ce.id, ce.uuid, ce.email, ce.start_time, ce.end_time, ce.event, ce.timezone, ce.created_at, ce.updated_at`
func (ds *Datastore) CreateOrUpdateCalendarEvent(
ctx context.Context,
uuid string,
uuidStr string,
email string,
startTime time.Time,
endTime time.Time,
@@ -23,11 +26,15 @@ func (ds *Datastore) CreateOrUpdateCalendarEvent(
hostID uint,
webhookStatus fleet.CalendarWebhookStatus,
) (*fleet.CalendarEvent, error) {
UUID, err := uuid.Parse(uuidStr)
if err != nil {
return nil, ctxerr.Wrap(ctx, err, "invalid uuid")
}
var id int64
if err := ds.withRetryTxx(ctx, func(tx sqlx.ExtContext) error {
const calendarEventsQuery = `
INSERT INTO calendar_events (
uuid,
uuid_bin,
email,
start_time,
end_time,
@@ -35,7 +42,7 @@ func (ds *Datastore) CreateOrUpdateCalendarEvent(
timezone
) VALUES (?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
uuid = VALUES(uuid),
uuid_bin = VALUES(uuid_bin),
start_time = VALUES(start_time),
end_time = VALUES(end_time),
event = VALUES(event),
@@ -45,7 +52,7 @@ func (ds *Datastore) CreateOrUpdateCalendarEvent(
result, err := tx.ExecContext(
ctx,
calendarEventsQuery,
uuid,
UUID[:],
email,
startTime,
endTime,
@@ -98,9 +105,7 @@ func (ds *Datastore) CreateOrUpdateCalendarEvent(
}
func getCalendarEventByID(ctx context.Context, q sqlx.QueryerContext, id uint) (*fleet.CalendarEvent, error) {
const calendarEventsQuery = `
SELECT * FROM calendar_events WHERE id = ?;
`
const calendarEventsQuery = "SELECT " + calendarEventCols + " FROM calendar_events ce WHERE id = ?"
var calendarEvent fleet.CalendarEvent
err := sqlx.GetContext(ctx, q, &calendarEvent, calendarEventsQuery, id)
if err != nil {
@@ -113,9 +118,7 @@ func getCalendarEventByID(ctx context.Context, q sqlx.QueryerContext, id uint) (
}
func (ds *Datastore) GetCalendarEvent(ctx context.Context, email string) (*fleet.CalendarEvent, error) {
const calendarEventsQuery = `
SELECT * FROM calendar_events WHERE email = ?;
`
const calendarEventsQuery = "SELECT " + calendarEventCols + " FROM calendar_events ce WHERE email = ?"
var calendarEvent fleet.CalendarEvent
err := sqlx.GetContext(ctx, ds.reader(ctx), &calendarEvent, calendarEventsQuery, email)
if err != nil {
@@ -127,29 +130,37 @@ func (ds *Datastore) GetCalendarEvent(ctx context.Context, email string) (*fleet
return &calendarEvent, nil
}
func (ds *Datastore) GetCalendarEventDetailsByUUID(ctx context.Context, uuid string) (*fleet.CalendarEventDetails, error) {
func (ds *Datastore) GetCalendarEventDetailsByUUID(ctx context.Context, uuidStr string) (*fleet.CalendarEventDetails, error) {
UUID, err := uuid.Parse(uuidStr)
if err != nil {
return nil, ctxerr.Wrap(ctx, err, "invalid uuid")
}
const calendarEventsByUUIDQuery = `
SELECT ce.*, h.team_id as team_id, h.id as host_id FROM calendar_events ce
SELECT ` + calendarEventCols + `, h.team_id as team_id, h.id as host_id FROM calendar_events ce
LEFT JOIN host_calendar_events hce ON hce.calendar_event_id = ce.id
LEFT JOIN hosts h ON h.id = hce.host_id
WHERE ce.uuid = ?;
WHERE ce.uuid_bin = ?;
`
var calendarEvent fleet.CalendarEventDetails
err := sqlx.GetContext(ctx, ds.reader(ctx), &calendarEvent, calendarEventsByUUIDQuery, uuid)
err = sqlx.GetContext(ctx, ds.reader(ctx), &calendarEvent, calendarEventsByUUIDQuery, UUID[:])
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, ctxerr.Wrap(ctx, notFound("CalendarEvent").WithMessage(fmt.Sprintf("uuid: %s", uuid)))
return nil, ctxerr.Wrap(ctx, notFound("CalendarEvent").WithMessage(fmt.Sprintf("uuid: %s", UUID.String())))
}
return nil, ctxerr.Wrap(ctx, err, "get calendar event")
}
return &calendarEvent, nil
}
func (ds *Datastore) UpdateCalendarEvent(ctx context.Context, calendarEventID uint, uuid string, startTime time.Time, endTime time.Time,
func (ds *Datastore) UpdateCalendarEvent(ctx context.Context, calendarEventID uint, uuidStr string, startTime time.Time, endTime time.Time,
data []byte, timeZone string) error {
UUID, err := uuid.Parse(uuidStr)
if err != nil {
return ctxerr.Wrap(ctx, err, "invalid uuid")
}
const calendarEventsQuery = `
UPDATE calendar_events SET
uuid = ?,
uuid_bin = ?,
start_time = ?,
end_time = ?,
event = ?,
@@ -157,7 +168,7 @@ func (ds *Datastore) UpdateCalendarEvent(ctx context.Context, calendarEventID ui
updated_at = CURRENT_TIMESTAMP
WHERE id = ?;
`
if _, err := ds.writer(ctx).ExecContext(ctx, calendarEventsQuery, uuid, startTime, endTime, data, timeZone,
if _, err := ds.writer(ctx).ExecContext(ctx, calendarEventsQuery, UUID[:], startTime, endTime, data, timeZone,
calendarEventID); err != nil {
return ctxerr.Wrap(ctx, err, "update calendar event")
}
@@ -186,7 +197,7 @@ func (ds *Datastore) GetHostCalendarEvent(ctx context.Context, hostID uint) (*fl
return nil, nil, ctxerr.Wrap(ctx, err, "get host calendar event")
}
const calendarEventsQuery = `
SELECT * FROM calendar_events WHERE id = ?
SELECT ` + calendarEventCols + ` FROM calendar_events ce WHERE id = ?
`
var calendarEvent fleet.CalendarEvent
if err := sqlx.GetContext(ctx, ds.reader(ctx), &calendarEvent, calendarEventsQuery, hostCalendarEvent.CalendarEventID); err != nil {
@@ -200,7 +211,7 @@ func (ds *Datastore) GetHostCalendarEvent(ctx context.Context, hostID uint) (*fl
func (ds *Datastore) GetHostCalendarEventByEmail(ctx context.Context, email string) (*fleet.HostCalendarEvent, *fleet.CalendarEvent, error) {
const calendarEventsQuery = `
SELECT * FROM calendar_events WHERE email = ?
SELECT ` + calendarEventCols + ` FROM calendar_events ce WHERE email = ?
`
var calendarEvent fleet.CalendarEvent
if err := sqlx.GetContext(ctx, ds.reader(ctx), &calendarEvent, calendarEventsQuery, email); err != nil {
@@ -236,7 +247,7 @@ func (ds *Datastore) UpdateHostCalendarWebhookStatus(ctx context.Context, hostID
func (ds *Datastore) ListCalendarEvents(ctx context.Context, teamID *uint) ([]*fleet.CalendarEvent, error) {
calendarEventsQuery := `
SELECT ce.* FROM calendar_events ce
SELECT ` + calendarEventCols + ` FROM calendar_events ce
`
var args []interface{}
@@ -258,7 +269,7 @@ func (ds *Datastore) ListCalendarEvents(ctx context.Context, teamID *uint) ([]*f
func (ds *Datastore) ListOutOfDateCalendarEvents(ctx context.Context, t time.Time) ([]*fleet.CalendarEvent, error) {
calendarEventsQuery := `
SELECT ce.* FROM calendar_events ce WHERE updated_at < ?
SELECT ` + calendarEventCols + ` FROM calendar_events ce WHERE updated_at < ?
`
var calendarEvents []*fleet.CalendarEvent
if err := sqlx.SelectContext(ctx, ds.reader(ctx), &calendarEvents, calendarEventsQuery, t); err != nil {
@@ -4,6 +4,7 @@ import (
"context"
"github.com/google/uuid"
"github.com/stretchr/testify/assert"
"strings"
"testing"
"time"
@@ -79,7 +80,7 @@ func testUpdateCalendarEvent(t *testing.T, ds *Datastore) {
eventDetails, err := ds.GetCalendarEventDetailsByUUID(ctx, eventUUIDNew)
require.NoError(t, err)
assert.Equal(t, eventUUIDNew, eventDetails.UUID)
assert.Equal(t, strings.ToUpper(eventUUIDNew), eventDetails.UUID)
assert.Equal(t, *calendarEvent, eventDetails.CalendarEvent)
assert.Equal(t, host.ID, eventDetails.HostID)
assert.Nil(t, eventDetails.TeamID)
@@ -21,7 +21,7 @@ func TestUp_20240707134035(t *testing.T) {
// Apply current migration.
applyNext(t, db)
// check that it's NULL
// check that UUID is not NULL
const selectUUIDStmt = `SELECT uuid FROM calendar_events WHERE id = ?`
var uuid1, uuid2 string
err := db.Get(&uuid1, selectUUIDStmt, event1ID)
@@ -0,0 +1,53 @@
package tables
import (
"database/sql"
"fmt"
)
func init() {
MigrationClient.AddMigration(Up_20240709132642, Down_20240709132642)
}
func Up_20240709132642(tx *sql.Tx) error {
// Implementation based on: https://dev.mysql.com/blog-archive/storing-uuid-values-in-mysql-tables/
if _, err := tx.Exec(`ALTER TABLE calendar_events ADD COLUMN uuid_bin BINARY(16) NOT NULL`); err != nil {
return fmt.Errorf("failed to add `uuid_bin` column to `calendar_events` table: %w", err)
}
// Convert existing UUIDs to binary format
if _, err := tx.Exec(`UPDATE calendar_events SET uuid_bin = UNHEX(REPLACE(uuid,'-','')), updated_at = updated_at`); err != nil {
return fmt.Errorf("failed to convert UUIDs to binary form for existing calendar events: %w", err)
}
// Add unique constraint to uuid_bin column
if _, err := tx.Exec(`ALTER TABLE calendar_events ADD CONSTRAINT idx_calendar_events_uuid_bin_unique UNIQUE (uuid_bin)`); err != nil {
return fmt.Errorf("failed to add unique constraint to `uuid_bin` column in `calendar_events` table: %w", err)
}
// Drop existing uuid column
if _, err := tx.Exec(`ALTER TABLE calendar_events DROP COLUMN uuid`); err != nil {
return fmt.Errorf("failed to drop `uuid` column from `calendar_events` table: %w", err)
}
// Add a new GENERATED uuid column
if _, err := tx.Exec(
`ALTER TABLE calendar_events ADD COLUMN uuid VARCHAR(36) COLLATE utf8mb4_unicode_ci GENERATED ALWAYS AS (
(INSERT(
INSERT(
INSERT(
INSERT(hex(uuid_bin),9,0,'-'),
14,0,'-'),
19,0,'-'),
24,0,'-')
)) VIRTUAL`); err != nil {
return fmt.Errorf("failed to add `uuid` column to `calendar_events` table: %w", err)
}
return nil
}
func Down_20240709132642(_ *sql.Tx) error {
return nil
}
@@ -0,0 +1,49 @@
package tables
import (
"github.com/google/uuid"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"strings"
"testing"
"time"
)
func TestUp_20240709132642(t *testing.T) {
db := applyUpToPrev(t)
testUUID := strings.ToUpper(uuid.New().String())
startTime := time.Now().UTC()
endTime := time.Now().UTC().Add(30 * time.Minute)
data := []byte("{\"foo\": \"bar\"}")
const insertStmtUUID = `INSERT INTO calendar_events (email, start_time, end_time, event, uuid) VALUES (?, ?, ?, ?, ?)`
eventID := execNoErrLastID(t, db, insertStmtUUID, "bob@example.com", startTime, endTime, data, testUUID)
applyNext(t, db)
// Check that uuid and uuid_bin are correct
const selectUUIDStmt = `SELECT uuid, uuid_bin FROM calendar_events WHERE id = ?`
type event struct {
UUID string `db:"uuid"`
UUIDBin []byte `db:"uuid_bin"`
}
var e event
err := db.Get(&e, selectUUIDStmt, eventID)
require.NoError(t, err)
assert.Equal(t, testUUID, e.UUID)
uuidFromBytes, err := uuid.FromBytes(e.UUIDBin)
require.NoError(t, err)
assert.Equal(t, uuid.MustParse(testUUID), uuidFromBytes)
// Try to use the same uuid again
const insertStmtUUIDBin = `INSERT INTO calendar_events (email, start_time, end_time, event, uuid_bin) VALUES (?, ?, ?, ?, ?)`
_, err = db.Exec(insertStmtUUIDBin, "alice@example.com", startTime, endTime, data, e.UUIDBin)
assert.Error(t, err)
// Insert a new event with a new UUID
uuidBin := uuid.New()
eventID = execNoErrLastID(t, db, insertStmtUUIDBin, "jane@example.com", startTime, endTime, data, uuidBin[:])
err = db.Get(&e, selectUUIDStmt, eventID)
require.NoError(t, err)
assert.Equal(t, uuidBin[:], e.UUIDBin)
assert.Equal(t, strings.ToUpper(uuidBin.String()), e.UUID)
}
+3 -3
View File
@@ -3650,11 +3650,11 @@ func testGetTeamHostsPolicyMemberships(t *testing.T, ds *Datastore) {
//
tZ := "America/Argentina/Buenos_Aires"
now := time.Now()
eventUUID1 := "event-uuid"
eventUUID1 := uuid.New().String()
_, err = ds.CreateOrUpdateCalendarEvent(ctx, eventUUID1, "foo@example.com", now, now.Add(30*time.Minute), []byte(`{"foo": "bar"}`), tZ,
host1.ID, fleet.CalendarWebhookStatusPending)
require.NoError(t, err)
eventUUID2 := "event-uuid2"
eventUUID2 := uuid.New().String()
_, err = ds.CreateOrUpdateCalendarEvent(ctx, eventUUID2, "bar@example.com", now, now.Add(30*time.Minute), []byte(`{"foo": "bar"}`), tZ,
host6.ID, fleet.CalendarWebhookStatusPending)
require.NoError(t, err)
@@ -3740,7 +3740,7 @@ func testGetTeamHostsPolicyMemberships(t *testing.T, ds *Datastore) {
_, err = ds.CreateOrUpdateCalendarEvent(ctx, eventUUID1, "foo@example.com", now, now.Add(30*time.Minute), []byte(`{"foo": "bar"}`), tZ,
host2.ID, fleet.CalendarWebhookStatusPending)
require.NoError(t, err)
eventUUID3 := "event-uuid3"
eventUUID3 := uuid.New().String()
calendarEventHost3, err := ds.CreateOrUpdateCalendarEvent(ctx, eventUUID3, "zoo@example.com", now, now.Add(30*time.Minute),
[]byte(`{"foo": "bar"}`), tZ, host3.ID, fleet.CalendarWebhookStatusPending)
require.NoError(t, err)
File diff suppressed because one or more lines are too long
+19
View File
@@ -42,6 +42,25 @@ type UserCalendar interface {
Get(event *CalendarEvent, key string) (interface{}, error)
}
// Lock interface for managing distributed locks.
type Lock interface {
// AcquireLock attempts to acquire a lock with the given key. value is the value to set for the key, which is used to release the lock.
// expireMs is the time in milliseconds after which the lock is automatically released. expireMs=0 means a default expiration time is used.
// Returns true if the lock was acquired, false otherwise.
AcquireLock(ctx context.Context, key string, value string, expireMs uint64) (ok bool, err error)
// ReleaseLock attempts to release a lock with the given key and value. If key does not exist or value does not match, the lock is not released.
// Returns true if the lock was released, false otherwise.
ReleaseLock(ctx context.Context, key string, value string) (ok bool, err error)
// Get retrieves the value of the given key. If the key does not exist, nil is returned.
Get(ctx context.Context, key string) (*string, error)
// AddToSet adds the value to the set identified by the given key.
AddToSet(ctx context.Context, key string, value string) error
// RemoveFromSet removes the value from the set identified by the given key.
RemoveFromSet(ctx context.Context, key string, value string) error
// GetSet retrieves a slice of string values from the set identified by the given key.
GetSet(ctx context.Context, key string) ([]string, error)
}
type CalendarWebhookPayload struct {
Timestamp time.Time `json:"timestamp"`
HostID uint `json:"host_id"`
+8 -2
View File
@@ -16,13 +16,19 @@ import (
"github.com/go-kit/log/level"
)
type CalendarConfig struct {
const (
LockKeyPrefix = "calendar:lock:"
ReservedLockKeyPrefix = "calendar:reserved:"
QueueKey = "calendar:queue"
)
type Config struct {
config.CalendarConfig
fleet.GoogleCalendarIntegration
ServerURL string
}
func CreateUserCalendarFromConfig(ctx context.Context, config *CalendarConfig, logger kitlog.Logger) fleet.UserCalendar {
func CreateUserCalendarFromConfig(ctx context.Context, config *Config, logger kitlog.Logger) fleet.UserCalendar {
googleCalendarConfig := calendar.GoogleCalendarConfig{
Context: ctx,
IntegrationConfig: &config.GoogleCalendarIntegration,
+153 -13
View File
@@ -35,6 +35,8 @@ import (
"github.com/fleetdm/fleet/v4/server/mdm"
"github.com/fleetdm/fleet/v4/server/ptr"
"github.com/fleetdm/fleet/v4/server/pubsub"
commonCalendar "github.com/fleetdm/fleet/v4/server/service/calendar"
"github.com/fleetdm/fleet/v4/server/service/redis_lock"
"github.com/fleetdm/fleet/v4/server/service/schedule"
"github.com/fleetdm/fleet/v4/server/test"
"github.com/go-kit/log"
@@ -87,7 +89,8 @@ func (s *integrationEnterpriseTestSuite) SetupSuite() {
cronLog = kitlog.NewNopLogger()
}
calendarSchedule, err = cron.NewCalendarSchedule(
ctx, s.T().Name(), s.ds, config.CalendarConfig{Periodicity: 24 * time.Hour}, cronLog,
ctx, s.T().Name(), s.ds, redis_lock.NewLock(s.redisPool), config.CalendarConfig{Periodicity: 24 * time.Hour},
cronLog,
)
return calendarSchedule, err
}
@@ -10959,9 +10962,9 @@ func (s *integrationEnterpriseTestSuite) TestCalendarCallback() {
},
)
require.NoError(t, err)
team1Policy2, err := s.ds.NewTeamPolicy(
team1Policy2Calendar, err := s.ds.NewTeamPolicy(
ctx, team1.ID, nil, fleet.PolicyPayload{
Name: "team1Policy2",
Name: "team1Policy2Calendar",
Query: "SELECT 2;",
CalendarEventsEnabled: true,
},
@@ -11012,7 +11015,7 @@ func (s *integrationEnterpriseTestSuite) TestCalendarCallback() {
host1Team1,
map[uint]*bool{
team1Policy1Calendar.ID: ptr.Bool(false),
team1Policy2.ID: ptr.Bool(true),
team1Policy2Calendar.ID: ptr.Bool(true),
globalPolicy.ID: nil,
},
), http.StatusOK, &distributedResp)
@@ -11022,7 +11025,7 @@ func (s *integrationEnterpriseTestSuite) TestCalendarCallback() {
host2Team1,
map[uint]*bool{
team1Policy1Calendar.ID: ptr.Bool(true),
team1Policy2.ID: ptr.Bool(false),
team1Policy2Calendar.ID: ptr.Bool(false),
globalPolicy.ID: nil,
},
), http.StatusOK, &distributedResp)
@@ -11106,28 +11109,95 @@ func (s *integrationEnterpriseTestSuite) TestCalendarCallback() {
// Delete the event on the calendar
calendar.ClearMockEvents()
// This callback should recreate the event
// Grab the distributed lock for this event
distributedLock := redis_lock.NewLock(s.redisPool)
lockValue := uuid.New().String()
result, err := distributedLock.AcquireLock(ctx, commonCalendar.LockKeyPrefix+event.UUID, lockValue, 0)
require.NoError(t, err)
assert.NotEmpty(t, result)
// This callback should put the event processing in a queue for async processing. It does not start async
// processing because it assumes another server is handling this webhook, and that server will start
// async processing.
_ = s.DoRawWithHeaders("POST", "/api/v1/fleet/calendar/webhook/"+event.UUID, []byte(""), http.StatusOK, map[string]string{
"X-Goog-Channel-Id": details.ChannelID,
"X-Goog-Resource-State": "exists",
})
team1CalendarEvents, err = s.ds.ListCalendarEvents(ctx, &team1.ID)
uuids, err := distributedLock.GetSet(ctx, commonCalendar.QueueKey)
require.NoError(t, err)
require.Len(t, team1CalendarEvents, 1)
assert.ElementsMatch(t, []string{event.UUID}, uuids)
// The calendar should still be empty since event hasn't processed yet
assert.Zero(t, len(calendar.ListGoogleMockEvents()))
// We clear the queue
assert.NoError(t, distributedLock.RemoveFromSet(ctx, commonCalendar.QueueKey, event.UUID))
// We release the normal lock, but grab the reserve lock instead
ok, err := distributedLock.ReleaseLock(ctx, commonCalendar.LockKeyPrefix+event.UUID, lockValue)
require.NoError(t, err)
assert.True(t, ok)
result, err = distributedLock.AcquireLock(ctx, commonCalendar.ReservedLockKeyPrefix+event.UUID, lockValue, 0)
require.NoError(t, err)
assert.NotEmpty(t, result)
// This callback should put the event processing in a queue for async processing, AND start the async processing
_ = s.DoRawWithHeaders("POST", "/api/v1/fleet/calendar/webhook/"+event.UUID, []byte(""), http.StatusOK, map[string]string{
"X-Goog-Channel-Id": details.ChannelID,
"X-Goog-Resource-State": "exists",
})
uuids, err = distributedLock.GetSet(ctx, commonCalendar.QueueKey)
require.NoError(t, err)
assert.ElementsMatch(t, []string{event.UUID}, uuids)
// The calendar should still be empty since event hasn't processed yet
assert.Zero(t, len(calendar.ListGoogleMockEvents()))
// We grab the normal lock again.
lockValue2 := uuid.New().String()
result, err = distributedLock.AcquireLock(ctx, commonCalendar.LockKeyPrefix+event.UUID, lockValue2, 0)
require.NoError(t, err)
assert.NotEmpty(t, result)
// We release the reserve lock.
ok, err = distributedLock.ReleaseLock(ctx, commonCalendar.ReservedLockKeyPrefix+event.UUID, lockValue)
require.NoError(t, err)
assert.True(t, ok)
// We release the normal lock.
ok, err = distributedLock.ReleaseLock(ctx, commonCalendar.LockKeyPrefix+event.UUID, lockValue2)
require.NoError(t, err)
assert.True(t, ok)
done := make(chan struct{})
go func() {
for {
time.Sleep(100 * time.Millisecond)
team1CalendarEvents, err = s.ds.ListCalendarEvents(ctx, &team1.ID)
require.NoError(t, err)
require.Len(t, team1CalendarEvents, 1)
if event.UUID != team1CalendarEvents[0].UUID {
done <- struct{}{}
return
}
}
}()
select {
case <-done: // All good
case <-time.After(5 * time.Second):
t.Fatal("timeout waiting for calendar event processing")
}
eventRecreated := team1CalendarEvents[0]
assert.NotZero(t, eventRecreated.ID)
assert.Equal(t, user1Email, eventRecreated.Email)
assert.NotZero(t, eventRecreated.StartTime)
assert.NotZero(t, eventRecreated.EndTime)
assert.NotEmpty(t, eventRecreated.UUID)
assert.NotEqual(t, event.UUID, eventRecreated.UUID)
assert.NotEqual(t, event.StartTime, eventRecreated.StartTime)
assert.NotEqual(t, event.EndTime, eventRecreated.EndTime)
assert.Equal(t, 1, calendar.MockChannelsCount())
assert.Equal(t, 1, len(calendar.ListGoogleMockEvents()))
// The previous event UUID should not work anymore
_ = s.DoRawWithHeaders("POST", "/api/v1/fleet/calendar/webhook/"+event.UUID, []byte(""), http.StatusNotFound, map[string]string{
// The previous event UUID should not work anymore, but API returns OK because this is a common occurrence.
_ = s.DoRawWithHeaders("POST", "/api/v1/fleet/calendar/webhook/"+event.UUID, []byte(""), http.StatusOK, map[string]string{
"X-Goog-Channel-Id": details.ChannelID,
"X-Goog-Resource-State": "exists",
})
@@ -11171,6 +11241,75 @@ func (s *integrationEnterpriseTestSuite) TestCalendarCallback() {
assert.Equal(t, eventRecreated.EndTime, eventUpdated.EndTime)
assert.Equal(t, 1, calendar.MockChannelsCount())
// Update the time of the event again
events = calendar.ListGoogleMockEvents()
require.Len(t, events, 1)
for _, e := range events {
st, err := time.Parse(time.RFC3339, e.Start.DateTime)
require.NoError(t, err)
newStartTime := st.Add(5 * time.Minute).Format(time.RFC3339)
e.Start.DateTime = newStartTime
}
// Grab the lock
event = eventUpdated
lockValue = uuid.New().String()
result, err = distributedLock.AcquireLock(ctx, commonCalendar.LockKeyPrefix+event.UUID, lockValue, 0)
require.NoError(t, err)
assert.NotEmpty(t, result)
mysql.ExecAdhocSQL(t, s.ds, func(db sqlx.ExtContext) error {
// Update updated_at so the event gets updated (the event is updated regularly)
_, err := db.ExecContext(ctx,
`UPDATE calendar_events SET updated_at = DATE_SUB(CURRENT_TIMESTAMP, INTERVAL 25 HOUR) WHERE id = ?`, event.ID)
return err
})
// Trigger the calendar cron async. It should wait for the lock and set reserve lock.
go triggerAndWait(ctx, t, s.ds, s.calendarSchedule, 10*time.Second)
done = make(chan struct{})
go func() {
for {
time.Sleep(100 * time.Millisecond)
reserveLock, err := distributedLock.Get(ctx, commonCalendar.ReservedLockKeyPrefix+event.UUID)
require.NoError(t, err)
if reserveLock != nil {
done <- struct{}{}
return
}
}
}()
select {
case <-done: // All good
case <-time.After(5 * time.Second):
t.Fatal("timeout waiting for cron to set reserve lock")
}
// Release the normal lock
ok, err = distributedLock.ReleaseLock(ctx, commonCalendar.LockKeyPrefix+event.UUID, lockValue)
require.NoError(t, err)
assert.True(t, ok)
// Wait for the event to update
done = make(chan struct{})
go func() {
for {
time.Sleep(100 * time.Millisecond)
team1CalendarEvents, err = s.ds.ListCalendarEvents(ctx, &team1.ID)
require.NoError(t, err)
if len(team1CalendarEvents) == 1 && team1CalendarEvents[0].UUID == event.UUID &&
team1CalendarEvents[0].StartTime.After(event.StartTime) {
done <- struct{}{}
return
}
}
}()
select {
case <-done: // All good
case <-time.After(5 * time.Second):
t.Fatal("timeout waiting for event to update during cron")
}
// Delete the event on the calendar
calendar.ClearMockEvents()
@@ -11179,7 +11318,7 @@ func (s *integrationEnterpriseTestSuite) TestCalendarCallback() {
host1Team1,
map[uint]*bool{
team1Policy1Calendar.ID: ptr.Bool(true),
team1Policy2.ID: ptr.Bool(true),
team1Policy2Calendar.ID: ptr.Bool(true),
globalPolicy.ID: nil,
},
), http.StatusOK, &distributedResp)
@@ -11192,10 +11331,11 @@ func (s *integrationEnterpriseTestSuite) TestCalendarCallback() {
})
assert.Equal(t, 0, calendar.MockChannelsCount())
previousEvent := team1CalendarEvents[0]
team1CalendarEvents, err = s.ds.ListCalendarEvents(ctx, &team1.ID)
require.NoError(t, err)
require.Len(t, team1CalendarEvents, 1)
assert.Equal(t, eventUpdated, team1CalendarEvents[0])
assert.Equal(t, previousEvent, team1CalendarEvents[0])
// Trigger calendar should cleanup the events
triggerAndWait(ctx, t, s.ds, s.calendarSchedule, 5*time.Second)
+118
View File
@@ -0,0 +1,118 @@
package redis_lock
import (
"context"
"errors"
"github.com/fleetdm/fleet/v4/server/contexts/ctxerr"
"github.com/fleetdm/fleet/v4/server/datastore/redis"
"github.com/fleetdm/fleet/v4/server/fleet"
redigo "github.com/gomodule/redigo/redis"
)
// This package implements a distributed lock using Redis. The lock can be used
// to prevent multiple Fleet servers from accessing a shared resource.
const (
defaultExpireMs = 60 * 1000
)
type redisLock struct {
pool fleet.RedisPool
testPrefix string // for tests, the key prefix to use to avoid conflicts
}
func NewLock(pool fleet.RedisPool) fleet.Lock {
lock := &redisLock{
pool: pool,
}
return fleet.Lock(lock)
}
func (r *redisLock) AcquireLock(ctx context.Context, key string, value string, expireMs uint64) (ok bool, err error) {
conn := redis.ConfigureDoer(r.pool, r.pool.Get())
defer conn.Close()
if expireMs == 0 {
expireMs = defaultExpireMs
}
// Reference: https://redis.io/docs/latest/commands/set/
// NX -- Only set the key if it does not already exist.
result, err := redigo.String(conn.Do("SET", r.testPrefix+key, value, "NX", "PX", expireMs))
if err != nil && !errors.Is(err, redigo.ErrNil) {
return false, ctxerr.Wrap(ctx, err, "redis acquire lock")
}
return result != "", nil
}
func (r *redisLock) ReleaseLock(ctx context.Context, key string, value string) (ok bool, err error) {
conn := redis.ConfigureDoer(r.pool, r.pool.Get())
defer conn.Close()
const unlockScript = `
if redis.call("get", KEYS[1]) == ARGV[1] then
return redis.call("del", KEYS[1])
else
return 0
end
`
// Reference: https://redis.io/docs/latest/commands/set/
// Only release the lock if the value matches.
res, err := redigo.Int64(conn.Do("EVAL", unlockScript, 1, r.testPrefix+key, value))
if err != nil && !errors.Is(err, redigo.ErrNil) {
return false, ctxerr.Wrap(ctx, err, "redis release lock")
}
return res > 0, nil
}
func (r *redisLock) AddToSet(ctx context.Context, key string, value string) error {
conn := redis.ConfigureDoer(r.pool, r.pool.Get())
defer conn.Close()
// Reference: https://redis.io/docs/latest/commands/sadd/
_, err := conn.Do("SADD", r.testPrefix+key, value)
if err != nil {
return ctxerr.Wrap(ctx, err, "redis add to set")
}
return nil
}
func (r *redisLock) RemoveFromSet(ctx context.Context, key string, value string) error {
conn := redis.ConfigureDoer(r.pool, r.pool.Get())
defer conn.Close()
// Reference: https://redis.io/docs/latest/commands/srem/
_, err := conn.Do("SREM", r.testPrefix+key, value)
if err != nil {
return ctxerr.Wrap(ctx, err, "redis add to set")
}
return nil
}
func (r *redisLock) GetSet(ctx context.Context, key string) ([]string, error) {
conn := redis.ConfigureDoer(r.pool, r.pool.Get())
defer conn.Close()
// Reference: https://redis.io/docs/latest/commands/smembers/
members, err := redigo.Strings(conn.Do("SMEMBERS", r.testPrefix+key))
if err != nil && !errors.Is(err, redigo.ErrNil) {
return nil, ctxerr.Wrap(ctx, err, "redis get set members")
}
return members, nil
}
func (r *redisLock) Get(ctx context.Context, key string) (*string, error) {
conn := redis.ConfigureDoer(r.pool, r.pool.Get())
defer conn.Close()
res, err := redigo.String(conn.Do("GET", r.testPrefix+key))
if errors.Is(err, redigo.ErrNil) {
return nil, nil
}
if err != nil {
return nil, ctxerr.Wrap(ctx, err, "redis get")
}
return &res, nil
}
@@ -0,0 +1,145 @@
package redis_lock
import (
"context"
"github.com/fleetdm/fleet/v4/server/datastore/redis/redistest"
"github.com/fleetdm/fleet/v4/server/fleet"
"github.com/fleetdm/fleet/v4/server/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"testing"
"time"
)
func TestRedisLock(t *testing.T) {
for _, f := range []func(*testing.T, fleet.Lock){
testRedisAcquireLock,
testRedisSet,
} {
t.Run(test.FunctionName(f), func(t *testing.T) {
t.Run("standalone", func(t *testing.T) {
lock := setupRedis(t, false, false)
f(t, lock)
})
t.Run("cluster", func(t *testing.T) {
lock := setupRedis(t, true, true)
f(t, lock)
})
})
}
}
func setupRedis(t testing.TB, cluster, redir bool) fleet.Lock {
pool := redistest.SetupRedis(t, t.Name(), cluster, redir, true)
return NewLockTest(t, pool)
}
type TestName interface {
Name() string
}
// NewFailingTest creates a redis policy set for failing policies to be used
// only in tests.
func NewLockTest(t TestName, pool fleet.RedisPool) fleet.Lock {
lock := &redisLock{
pool: pool,
testPrefix: t.Name() + ":",
}
return fleet.Lock(lock)
}
func testRedisAcquireLock(t *testing.T, lock fleet.Lock) {
ctx := context.Background()
result, err := lock.AcquireLock(ctx, "test", "1", 0)
require.NoError(t, err)
assert.True(t, result)
// Try to acquire the same lock
result, err = lock.AcquireLock(ctx, "test", "1", 0)
assert.NoError(t, err)
assert.False(t, result)
// Try to release the lock with a wrong value
ok, err := lock.ReleaseLock(ctx, "test", "2")
require.NoError(t, err)
assert.False(t, ok)
// Try to release the lock with the wrong key
ok, err = lock.ReleaseLock(ctx, "bad", "1")
require.NoError(t, err)
assert.False(t, ok)
// Try to release the lock with the correct key/value
ok, err = lock.ReleaseLock(ctx, "test", "1")
require.NoError(t, err)
assert.True(t, ok)
// Acquire the lock again
result, err = lock.AcquireLock(ctx, "test", "1", 0)
require.NoError(t, err)
assert.True(t, result)
// Get lock
getResult, err := lock.Get(ctx, "test")
assert.NoError(t, err)
require.NotNil(t, getResult)
assert.Equal(t, "1", *getResult)
// Try to set lock with expiration
var expire uint64 = 10
result, err = lock.AcquireLock(ctx, "testE", "1", expire)
require.NoError(t, err)
assert.True(t, result)
// Try to acquire the same lock after waiting
duration := time.Duration(expire+1) * time.Millisecond
time.Sleep(duration)
result, err = lock.AcquireLock(ctx, "testE", "1", 0)
require.NoError(t, err)
assert.True(t, result)
// Get non-existent key
getResult, err = lock.Get(ctx, "testNonExistent")
assert.NoError(t, err)
assert.Nil(t, getResult)
}
func testRedisSet(t *testing.T, lock fleet.Lock) {
ctx := context.Background()
// Get a non-existent set
result, err := lock.GetSet(ctx, "missingSet")
assert.NoError(t, err)
assert.Empty(t, result)
// Add to a set
values := []string{"foo", "bar"}
err = lock.AddToSet(ctx, "testSet", values[0])
assert.NoError(t, err)
err = lock.AddToSet(ctx, "testSet", values[1])
assert.NoError(t, err)
// Get the set
result, err = lock.GetSet(ctx, "testSet")
assert.NoError(t, err)
assert.ElementsMatch(t, values, result)
// Remove from set
err = lock.RemoveFromSet(ctx, "testSet", values[0])
assert.NoError(t, err)
// Get the set
result, err = lock.GetSet(ctx, "testSet")
assert.NoError(t, err)
assert.Equal(t, []string{values[1]}, result)
// Remove from set
err = lock.RemoveFromSet(ctx, "testSet", values[1])
assert.NoError(t, err)
// Get the set
result, err = lock.GetSet(ctx, "testSet")
assert.NoError(t, err)
assert.Empty(t, result)
}
+4
View File
@@ -34,6 +34,7 @@ import (
"github.com/fleetdm/fleet/v4/server/ptr"
"github.com/fleetdm/fleet/v4/server/service/async"
"github.com/fleetdm/fleet/v4/server/service/mock"
"github.com/fleetdm/fleet/v4/server/service/redis_lock"
"github.com/fleetdm/fleet/v4/server/sso"
"github.com/fleetdm/fleet/v4/server/test"
kitlog "github.com/go-kit/log"
@@ -69,6 +70,7 @@ func newTestServiceWithConfig(t *testing.T, ds fleet.Datastore, fleetConfig conf
ssoStore sso.SessionStore
profMatcher fleet.ProfileMatcher
softwareInstallStore fleet.SoftwareInstallerStore
distributedLock fleet.Lock
)
if len(opts) > 0 {
if opts[0].Clock != nil {
@@ -95,6 +97,7 @@ func newTestServiceWithConfig(t *testing.T, ds fleet.Datastore, fleetConfig conf
if opts[0].Pool != nil {
ssoStore = sso.NewSessionStore(opts[0].Pool)
profMatcher = apple_mdm.NewProfileMatcher(opts[0].Pool)
distributedLock = redis_lock.NewLock(opts[0].Pool)
}
if opts[0].ProfileMatcher != nil {
profMatcher = opts[0].ProfileMatcher
@@ -194,6 +197,7 @@ func newTestServiceWithConfig(t *testing.T, ds fleet.Datastore, fleetConfig conf
ssoStore,
profMatcher,
softwareInstallStore,
distributedLock,
)
if err != nil {
panic(err)
+19 -1
View File
@@ -70,16 +70,20 @@ func main() {
log.Fatalf("Unable to create Calendar service: %v", err)
}
numberDeleted := 0
var maxResults int64 = 1000
pageToken := ""
now := time.Now()
for {
list, err := withRetry(
func() (any, error) {
return service.Events.List("primary").
EventTypes("default").
MaxResults(1000).
MaxResults(maxResults).
OrderBy("startTime").
SingleEvents(true).
ShowDeleted(false).
Q(eventTitle).
PageToken(pageToken).
Do()
},
)
@@ -89,7 +93,17 @@ func main() {
if len(list.(*calendar.Events).Items) == 0 {
break
}
foundNewEvents := false
for _, item := range list.(*calendar.Events).Items {
created, err := time.Parse(time.RFC3339, item.Created)
if err != nil {
log.Fatalf("Unable to parse event created time: %v", err)
}
if created.After(now) {
// Found events created after we started deleting events, so we should stop
foundNewEvents = true
continue // Skip this event but finish the loop to make sure we don't miss something
}
if item.Summary == eventTitle {
_, err := withRetry(
func() (any, error) {
@@ -105,6 +119,10 @@ func main() {
}
}
}
pageToken = list.(*calendar.Events).NextPageToken
if pageToken == "" || foundNewEvents {
break
}
}
log.Printf("DONE. Deleted %d events total for %s", numberDeleted, userEmail)
}(userEmail)
+19 -1
View File
@@ -81,16 +81,20 @@ func main() {
}
numberMoved := 0
var maxResults int64 = 1000
pageToken := ""
now := time.Now()
for {
list, err := withRetry(
func() (any, error) {
return service.Events.List("primary").EventTypes("default").
MaxResults(1000).
MaxResults(maxResults).
OrderBy("startTime").
SingleEvents(true).
ShowDeleted(false).
TimeMin(dateTimeEndStr).
Q(eventTitle).
PageToken(pageToken).
Do()
},
)
@@ -101,7 +105,17 @@ func main() {
if len(list.(*calendar.Events).Items) == 0 {
break
}
foundNewEvents := false
for _, item := range list.(*calendar.Events).Items {
created, err := time.Parse(time.RFC3339, item.Created)
if err != nil {
log.Fatalf("Unable to parse event created time: %v", err)
}
if created.After(now) {
// Found events created after we started moving events, so we should stop
foundNewEvents = true
continue // Skip this event but finish the loop to make sure we don't miss something
}
if item.Summary == eventTitle {
item.Start.DateTime = dateTime.Format(time.RFC3339)
item.End.DateTime = dateTime.Add(30 * time.Minute).Format(time.RFC3339)
@@ -120,6 +134,10 @@ func main() {
}
}
pageToken = list.(*calendar.Events).NextPageToken
if pageToken == "" || foundNewEvents {
break
}
}
log.Printf("DONE. Moved total %d events for %s", numberMoved, userEmail)
}(userEmail)