Fix MDM lock command race condition preventing duplicate PINs (#33001)
Fixes #31636. Prevents race condition where multiple concurrent lock commands could create multiple PINs, making hosts inaccessible.
This commit is contained in:
@@ -599,7 +599,7 @@ fleetctl mdm unlock --host=%s
|
||||
{appCfgMacMDM, "valid macos but pending", []string{"--host", macPending.host.UUID}, `Can't lock the host because it doesn't have MDM turned on.`},
|
||||
{appCfgAllMDM, "valid windows but pending unlock", []string{"--host", winEnrolledUP.host.UUID}, "Host has pending unlock request."},
|
||||
{appCfgAllMDM, "valid windows but pending lock", []string{"--host", winEnrolledLP.host.UUID}, "Host has pending lock request."},
|
||||
{appCfgAllMDM, "valid macos but pending lock", []string{"--host", macEnrolledLP.host.UUID}, "Host has pending lock request."},
|
||||
{appCfgAllMDM, "valid macos but pending lock", []string{"--host", macEnrolledLP.host.UUID}, ""},
|
||||
{appCfgAllMDM, "valid windows but pending wipe", []string{"--host", winEnrolledWP.host.UUID}, "Host has pending wipe request."},
|
||||
{appCfgAllMDM, "valid macos but pending wipe", []string{"--host", macEnrolledWP.host.UUID}, "Host has pending wipe request."},
|
||||
}
|
||||
@@ -1318,6 +1318,11 @@ func setupTestServer(t *testing.T) *mock.Store {
|
||||
return nil
|
||||
}
|
||||
|
||||
enqueuer.GetPendingLockCommandFunc = func(ctx context.Context, hostUUID string) (*mdm.Command, string, error) {
|
||||
// Return nil to indicate no pending lock command
|
||||
return nil, "", nil
|
||||
}
|
||||
|
||||
_, ds := testing_utils.RunServerWithMockedDS(t, &service.TestServerOpts{
|
||||
MDMStorage: enqueuer,
|
||||
MDMPusher: testing_utils.MockPusher{},
|
||||
|
||||
@@ -120,11 +120,20 @@ func (svc *Service) LockHost(ctx context.Context, hostID uint, viewPIN bool) (un
|
||||
if err != nil {
|
||||
return "", ctxerr.Wrap(ctx, err, "get host lock/wipe status")
|
||||
}
|
||||
|
||||
switch {
|
||||
case lockWipe.IsPendingLock():
|
||||
return "", ctxerr.Wrap(
|
||||
ctx, fleet.NewInvalidArgumentError("host_id", "Host has pending lock request. The host will lock when it comes online."),
|
||||
)
|
||||
// For macOS, we handle duplicate lock requests at the MDM commander level
|
||||
// by returning the existing PIN. For Windows/Linux, we need to prevent
|
||||
// duplicate script executions.
|
||||
if host.FleetPlatform() != "darwin" {
|
||||
return "", ctxerr.Wrap(
|
||||
ctx, fleet.NewInvalidArgumentError(
|
||||
"host_id", "Host has pending lock request. Host cannot be locked again until lock is complete.",
|
||||
),
|
||||
)
|
||||
}
|
||||
// For macOS, fall through to enqueueLockHostRequest which will handle the duplicate
|
||||
case lockWipe.IsPendingUnlock():
|
||||
return "", ctxerr.Wrap(
|
||||
ctx, fleet.NewInvalidArgumentError(
|
||||
|
||||
@@ -3,6 +3,7 @@ package mysql
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
@@ -23,6 +24,30 @@ import (
|
||||
nanomdm_log "github.com/micromdm/nanolib/log"
|
||||
)
|
||||
|
||||
// lockConflictError indicates a lock command already exists for the host
|
||||
type lockConflictError struct {
|
||||
hostUUID string
|
||||
}
|
||||
|
||||
func (e lockConflictError) Error() string {
|
||||
return "host already has a pending lock command"
|
||||
}
|
||||
|
||||
func (e lockConflictError) IsConflict() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// isConflict checks if an error implements the IsConflict() interface
|
||||
func isConflict(err error) bool {
|
||||
type conflictInterface interface {
|
||||
IsConflict() bool
|
||||
}
|
||||
if c, ok := err.(conflictInterface); ok {
|
||||
return c.IsConflict()
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// NanoMDMStorage wraps a *nanomdm_mysql.MySQLStorage and overrides further functionality.
|
||||
type NanoMDMStorage struct {
|
||||
*nanomdm_mysql.MySQLStorage
|
||||
@@ -114,12 +139,57 @@ func (s *NanoMDMStorage) StorePushCert(ctx context.Context, pemCert, pemKey []by
|
||||
return errors.New("please use fleet.Datastore to manage MDM assets")
|
||||
}
|
||||
|
||||
// GetPendingLockCommand returns the most recent unacknowledged DeviceLock command
|
||||
// for the given host, along with its unlock PIN.
|
||||
// Returns nil, "", nil if no pending lock command exists.
|
||||
func (s *NanoMDMStorage) GetPendingLockCommand(ctx context.Context, hostUUID string) (*mdm.Command, string, error) {
|
||||
query := `
|
||||
SELECT nc.command_uuid, nc.request_type, nc.command, hma.unlock_pin
|
||||
FROM nano_commands nc
|
||||
INNER JOIN host_mdm_actions hma ON hma.lock_ref = nc.command_uuid
|
||||
LEFT JOIN nano_command_results ncr ON ncr.command_uuid = nc.command_uuid
|
||||
INNER JOIN nano_enrollment_queue neq ON neq.command_uuid = nc.command_uuid
|
||||
WHERE neq.id = ?
|
||||
AND nc.request_type = 'DeviceLock'
|
||||
AND ncr.command_uuid IS NULL
|
||||
ORDER BY nc.created_at DESC
|
||||
LIMIT 1`
|
||||
|
||||
var result struct {
|
||||
CommandUUID string `db:"command_uuid"`
|
||||
RequestType string `db:"request_type"`
|
||||
Command []byte `db:"command"`
|
||||
UnlockPIN string `db:"unlock_pin"`
|
||||
}
|
||||
|
||||
err := sqlx.GetContext(ctx, s.db, &result, query, hostUUID)
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, "", nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, "", ctxerr.Wrap(ctx, err, "getting pending lock command")
|
||||
}
|
||||
|
||||
cmd := &mdm.Command{
|
||||
CommandUUID: result.CommandUUID,
|
||||
Command: struct {
|
||||
RequestType string
|
||||
}{
|
||||
RequestType: result.RequestType,
|
||||
},
|
||||
Raw: result.Command,
|
||||
}
|
||||
|
||||
return cmd, result.UnlockPIN, nil
|
||||
}
|
||||
|
||||
// EnqueueDeviceLockCommand enqueues a DeviceLock command for the given host.
|
||||
//
|
||||
// A few implementation details:
|
||||
// - It can only be called for a single hosts, to ensure we don't use the same
|
||||
// pin for multiple hosts.
|
||||
// - The method performs fleet-specific actions after the command is enqueued.
|
||||
// - It will fail with a ConflictError if a lock command already exists.
|
||||
func (s *NanoMDMStorage) EnqueueDeviceLockCommand(
|
||||
ctx context.Context,
|
||||
host *fleet.Host,
|
||||
@@ -127,10 +197,29 @@ func (s *NanoMDMStorage) EnqueueDeviceLockCommand(
|
||||
pin string,
|
||||
) error {
|
||||
return common_mysql.WithRetryTxx(ctx, s.db, func(tx sqlx.ExtContext) error {
|
||||
// check if a lock already exists using SELECT FOR UPDATE to prevent a race
|
||||
var existingLockRef *string
|
||||
err := sqlx.GetContext(ctx, tx, &existingLockRef,
|
||||
`SELECT lock_ref FROM host_mdm_actions WHERE host_id = ? FOR UPDATE`,
|
||||
host.ID)
|
||||
|
||||
// If we got a row and it has a lock_ref, fail with conflict
|
||||
if err == nil && existingLockRef != nil && *existingLockRef != "" {
|
||||
// A lock command already exists, don't overwrite
|
||||
return lockConflictError{hostUUID: host.UUID}
|
||||
}
|
||||
|
||||
// If the row doesn't exist, that's OK, we'll insert it
|
||||
if err != nil && err != sql.ErrNoRows {
|
||||
return ctxerr.Wrap(ctx, err, "checking for existing lock")
|
||||
}
|
||||
|
||||
// Now enqueue the command
|
||||
if err := enqueueCommandDB(ctx, tx, []string{host.UUID}, cmd); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Insert or update the host_mdm_actions row
|
||||
stmt := `
|
||||
INSERT INTO host_mdm_actions (
|
||||
host_id,
|
||||
|
||||
@@ -2,6 +2,8 @@ package mysql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -9,6 +11,8 @@ import (
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/nanomdm/mdm"
|
||||
"github.com/fleetdm/fleet/v4/server/ptr"
|
||||
"github.com/fleetdm/fleet/v4/server/test"
|
||||
"github.com/go-kit/log"
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
@@ -19,6 +23,8 @@ func TestNanoMDMStorage(t *testing.T) {
|
||||
fn func(t *testing.T, ds *Datastore)
|
||||
}{
|
||||
{"TestEnqueueDeviceLockCommand", testEnqueueDeviceLockCommand},
|
||||
{"TestGetPendingLockCommand", testGetPendingLockCommand},
|
||||
{"TestEnqueueDeviceLockCommandRaceCondition", testEnqueueDeviceLockCommandRaceCondition},
|
||||
}
|
||||
|
||||
for _, c := range cases {
|
||||
@@ -83,3 +89,228 @@ func testEnqueueDeviceLockCommand(t *testing.T, ds *Datastore) {
|
||||
require.Equal(t, "cmd-uuid", status.LockMDMCommand.CommandUUID)
|
||||
require.Equal(t, "123456", status.UnlockPIN)
|
||||
}
|
||||
|
||||
func testGetPendingLockCommand(t *testing.T, ds *Datastore) {
|
||||
ctx := context.Background()
|
||||
ns, err := ds.NewMDMAppleMDMStorage()
|
||||
require.NoError(t, err)
|
||||
|
||||
host, err := ds.NewHost(ctx, &fleet.Host{
|
||||
Hostname: "test-host2-name",
|
||||
OsqueryHostID: ptr.String("1338"),
|
||||
NodeKey: ptr.String("1338"),
|
||||
UUID: "test-uuid-2",
|
||||
TeamID: nil,
|
||||
Platform: "darwin",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
nanoEnroll(t, ds, host, false)
|
||||
|
||||
// Test 1: No pending commands should return nil
|
||||
cmd, pin, err := ns.GetPendingLockCommand(ctx, host.UUID)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, cmd)
|
||||
require.Empty(t, pin)
|
||||
|
||||
// Test 2: Enqueue a lock command
|
||||
lockCmd := &mdm.Command{}
|
||||
lockCmd.CommandUUID = "lock-cmd-uuid"
|
||||
lockCmd.Command.RequestType = "DeviceLock"
|
||||
lockCmd.Raw = []byte("<?xml")
|
||||
|
||||
err = ns.EnqueueDeviceLockCommand(ctx, host, lockCmd, "654321")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Test 3: Should find the pending command
|
||||
cmd, pin, err = ns.GetPendingLockCommand(ctx, host.UUID)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, cmd)
|
||||
require.Equal(t, "lock-cmd-uuid", cmd.CommandUUID)
|
||||
require.Equal(t, "DeviceLock", cmd.Command.RequestType)
|
||||
require.Equal(t, "654321", pin)
|
||||
|
||||
// Test 4: Multiple commands should fail due to conflict
|
||||
lockCmd2 := &mdm.Command{}
|
||||
lockCmd2.CommandUUID = "lock-cmd-uuid-2"
|
||||
lockCmd2.Command.RequestType = "DeviceLock"
|
||||
lockCmd2.Raw = []byte("<?xml2")
|
||||
|
||||
// This should fail with conflict error since a lock is already pending
|
||||
err = ns.EnqueueDeviceLockCommand(ctx, host, lockCmd2, "111111")
|
||||
require.Error(t, err)
|
||||
require.True(t, isConflict(err), "Should get conflict error for duplicate lock")
|
||||
|
||||
// The pending command should still be the original one
|
||||
cmd, pin, err = ns.GetPendingLockCommand(ctx, host.UUID)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, cmd)
|
||||
require.Equal(t, "lock-cmd-uuid", cmd.CommandUUID)
|
||||
require.Equal(t, "654321", pin)
|
||||
|
||||
// Test 5: After acknowledgment, should not find the command
|
||||
// First acknowledge the existing command
|
||||
_, err = ds.writer(ctx).ExecContext(ctx, `
|
||||
INSERT INTO nano_command_results (id, command_uuid, status, result)
|
||||
VALUES (?, ?, 'Acknowledged', '<?xml version="1.0"?><plist></plist>')`,
|
||||
host.UUID, "lock-cmd-uuid")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Now no pending command should exist
|
||||
cmd, pin, err = ns.GetPendingLockCommand(ctx, host.UUID)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, cmd)
|
||||
require.Empty(t, pin)
|
||||
|
||||
// Test 6: After acknowledgment, the lock_ref still exists in host_mdm_actions
|
||||
// This is expected behavior - the device remains locked until manually unlocked
|
||||
// Therefore, attempting to create a new lock command should still fail
|
||||
lockCmd3 := &mdm.Command{}
|
||||
lockCmd3.CommandUUID = "lock-cmd-uuid-3"
|
||||
lockCmd3.Command.RequestType = "DeviceLock"
|
||||
lockCmd3.Raw = []byte("<?xml3")
|
||||
|
||||
err = ns.EnqueueDeviceLockCommand(ctx, host, lockCmd3, "222222")
|
||||
require.Error(t, err)
|
||||
require.True(t, isConflict(err), "Should still get conflict error since device is locked")
|
||||
|
||||
// No pending command should exist since the previous was acknowledged
|
||||
cmd, pin, err = ns.GetPendingLockCommand(ctx, host.UUID)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, cmd)
|
||||
require.Empty(t, pin)
|
||||
}
|
||||
|
||||
func testEnqueueDeviceLockCommandRaceCondition(t *testing.T, ds *Datastore) {
|
||||
ctx := context.Background()
|
||||
|
||||
// Create a test host
|
||||
host, err := ds.NewHost(ctx, &fleet.Host{
|
||||
UUID: "test-host-race-" + uuid.NewString(),
|
||||
Platform: "darwin",
|
||||
OsqueryHostID: ptr.String("test-osquery-id"),
|
||||
NodeKey: ptr.String("test-node-key"),
|
||||
Hostname: "test-host.local",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Enable MDM for the host
|
||||
err = ds.SetOrUpdateMDMData(ctx, host.ID, false, true, "https://test.local", false, "test-ref", "", false)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create nano_devices record first
|
||||
deviceID := "device-" + host.UUID
|
||||
_, err = ds.writer(ctx).Exec(`
|
||||
INSERT INTO nano_devices (id, authenticate, token_update) VALUES (?, 'Authenticate', 0)`, deviceID)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create nano_enrollments record (required for MDM commands)
|
||||
_, err = ds.writer(ctx).Exec(`
|
||||
INSERT INTO nano_enrollments (id, device_id, type, topic, push_magic, token_hex, last_seen_at)
|
||||
VALUES (?, ?, 'Device', 'com.apple.mgmt.test', 'test-magic', 'deadbeef', NOW())`,
|
||||
host.UUID, deviceID)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create NanoMDMStorage
|
||||
storage := &NanoMDMStorage{
|
||||
db: ds.writer(ctx),
|
||||
logger: log.NewNopLogger(),
|
||||
ds: ds,
|
||||
}
|
||||
|
||||
// Number of concurrent lock attempts
|
||||
numGoroutines := 20
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(numGoroutines)
|
||||
|
||||
// Track successful locks
|
||||
successCount := 0
|
||||
var successMu sync.Mutex
|
||||
|
||||
// Track conflict errors
|
||||
conflictCount := 0
|
||||
|
||||
// Collect all PINs that were generated
|
||||
var pins []string
|
||||
|
||||
// Barrier to ensure all goroutines start at the same time
|
||||
barrier := make(chan struct{})
|
||||
|
||||
for i := 0; i < numGoroutines; i++ {
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
|
||||
// Wait for barrier
|
||||
<-barrier
|
||||
|
||||
// Create a unique command for this goroutine
|
||||
cmdUUID := fmt.Sprintf("test-lock-%03d", idx)
|
||||
pin := fmt.Sprintf("%06d", 100000+idx) // Unique PIN for each request
|
||||
|
||||
cmd := &mdm.Command{
|
||||
CommandUUID: cmdUUID,
|
||||
Command: struct {
|
||||
RequestType string
|
||||
}{
|
||||
RequestType: "DeviceLock",
|
||||
},
|
||||
Raw: []byte(fmt.Sprintf(`<?xml version="1.0"?><plist><dict><key>PIN</key><string>%s</string></dict></plist>`, pin)),
|
||||
}
|
||||
|
||||
// Try to enqueue the lock command
|
||||
err := storage.EnqueueDeviceLockCommand(ctx, host, cmd, pin)
|
||||
|
||||
switch {
|
||||
case err == nil:
|
||||
successMu.Lock()
|
||||
successCount++
|
||||
pins = append(pins, pin)
|
||||
successMu.Unlock()
|
||||
case isConflict(err):
|
||||
successMu.Lock()
|
||||
conflictCount++
|
||||
successMu.Unlock()
|
||||
default:
|
||||
// Unexpected error
|
||||
t.Logf("Request %d got unexpected error: %v", idx, err)
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
|
||||
// Release all goroutines at once
|
||||
close(barrier)
|
||||
|
||||
// Wait for all to complete
|
||||
wg.Wait()
|
||||
|
||||
// Check the database state
|
||||
|
||||
// 1. Count how many DeviceLock commands were created
|
||||
var commandCount int
|
||||
err = ds.writer(ctx).Get(&commandCount,
|
||||
`SELECT COUNT(*) FROM nano_commands WHERE command_uuid LIKE 'test-lock-%'`)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 2. Check what's stored in host_mdm_actions
|
||||
var storedPIN string
|
||||
var lockRef string
|
||||
err = ds.writer(ctx).QueryRow(
|
||||
`SELECT COALESCE(unlock_pin, ''), COALESCE(lock_ref, '') FROM host_mdm_actions WHERE host_id = ?`,
|
||||
host.ID).Scan(&storedPIN, &lockRef)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Log the results
|
||||
t.Logf("===== RACE CONDITION TEST RESULTS =====")
|
||||
t.Logf("Concurrent requests sent: %d", numGoroutines)
|
||||
t.Logf("Successful lock commands: %d", successCount)
|
||||
t.Logf("Conflict errors: %d", conflictCount)
|
||||
t.Logf("Commands in nano_commands table: %d", commandCount)
|
||||
t.Logf("Final PIN stored in database: %s", storedPIN)
|
||||
t.Logf("Final lock_ref in database: %s", lockRef)
|
||||
|
||||
// Assertions - only one lock should succeed
|
||||
require.Equal(t, 1, successCount, "Only one lock command should succeed")
|
||||
require.Equal(t, numGoroutines-1, conflictCount, "All other requests should get conflict error")
|
||||
require.Equal(t, 1, commandCount, "Only one command should be in nano_commands table")
|
||||
require.Len(t, pins, 1, "Only one PIN should be generated")
|
||||
require.Equal(t, pins[0], storedPIN, "Stored PIN should match the successful request")
|
||||
}
|
||||
|
||||
@@ -2387,6 +2387,7 @@ type AndroidDatastore interface {
|
||||
type MDMAppleStore interface {
|
||||
storage.AllStorage
|
||||
MDMAssetRetriever
|
||||
GetPendingLockCommand(ctx context.Context, hostUUID string) (*mdm.Command, string, error)
|
||||
EnqueueDeviceLockCommand(ctx context.Context, host *Host, cmd *mdm.Command, pin string) error
|
||||
EnqueueDeviceWipeCommand(ctx context.Context, host *Host, cmd *mdm.Command) error
|
||||
}
|
||||
|
||||
@@ -684,6 +684,11 @@ func (e ConflictError) StatusCode() int {
|
||||
return http.StatusConflict
|
||||
}
|
||||
|
||||
// IsConflict implements the conflict interface for middleware compatibility
|
||||
func (e ConflictError) IsConflict() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// Errorer interface is implemented by response structs to encode business logic errors
|
||||
type Errorer interface {
|
||||
Error() error
|
||||
|
||||
@@ -101,6 +101,21 @@ func (svc *MDMAppleCommander) RemoveProfile(ctx context.Context, hostUUIDs []str
|
||||
}
|
||||
|
||||
func (svc *MDMAppleCommander) DeviceLock(ctx context.Context, host *fleet.Host, uuid string) (unlockPIN string, err error) {
|
||||
// Check for existing pending lock command first
|
||||
existingCmd, existingPIN, err := svc.storage.GetPendingLockCommand(ctx, host.UUID)
|
||||
if err != nil {
|
||||
return "", ctxerr.Wrap(ctx, err, "checking for pending lock command")
|
||||
}
|
||||
|
||||
// If a pending lock command exists, just send a push notification and return the existing PIN
|
||||
if existingCmd != nil {
|
||||
if err := svc.SendNotifications(ctx, []string{host.UUID}); err != nil {
|
||||
return "", ctxerr.Wrap(ctx, err, "sending notifications for existing DeviceLock")
|
||||
}
|
||||
return existingPIN, nil
|
||||
}
|
||||
|
||||
// No pending lock, create a new one
|
||||
unlockPIN = GenerateRandomPin(6)
|
||||
raw := fmt.Sprintf(`<?xml version="1.0" encoding="UTF-8"?>
|
||||
<!DOCTYPE plist PUBLIC "-//Apple//DTD PLIST 1.0//EN" "http://www.apple.com/DTDs/PropertyList-1.0.dtd">
|
||||
@@ -125,6 +140,29 @@ func (svc *MDMAppleCommander) DeviceLock(ctx context.Context, host *fleet.Host,
|
||||
}
|
||||
|
||||
if err := svc.storage.EnqueueDeviceLockCommand(ctx, host, cmd, unlockPIN); err != nil {
|
||||
// Check if another request just created a lock
|
||||
type conflictInterface interface {
|
||||
IsConflict() bool
|
||||
}
|
||||
if c, ok := err.(conflictInterface); ok && c.IsConflict() {
|
||||
// Another goroutine won the race, fetch the command that was created
|
||||
existingCmd, existingPIN, err := svc.storage.GetPendingLockCommand(ctx, host.UUID)
|
||||
if err != nil {
|
||||
return "", ctxerr.Wrap(ctx, err, "getting existing lock after race condition")
|
||||
}
|
||||
if existingCmd != nil {
|
||||
// Send push notification for the existing command and return its PIN
|
||||
if pushErr := svc.SendNotifications(ctx, []string{host.UUID}); pushErr != nil {
|
||||
// Log the push error but still return the PIN since the command exists
|
||||
// The push can be retried on subsequent requests
|
||||
ctxerr.Handle(ctx, ctxerr.Wrap(ctx, pushErr, "failed to send push notification after lock race"))
|
||||
return existingPIN, nil
|
||||
}
|
||||
return existingPIN, nil
|
||||
}
|
||||
// This shouldn't happen, but if we can't find the command, return the original error
|
||||
return "", ctxerr.Wrap(ctx, err, "lock command conflict but no existing command found")
|
||||
}
|
||||
return "", ctxerr.Wrap(ctx, err, "enqueuing for DeviceLock")
|
||||
}
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/fleetdm/fleet/v4/server/fleet"
|
||||
@@ -24,6 +25,19 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// mockConflictError is used in tests to simulate a conflict error
|
||||
type mockConflictError struct {
|
||||
msg string
|
||||
}
|
||||
|
||||
func (e *mockConflictError) Error() string {
|
||||
return e.msg
|
||||
}
|
||||
|
||||
func (e *mockConflictError) IsConflict() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func TestMDMAppleCommander(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mdmStorage := &mdmmock.MDMAppleStore{}
|
||||
@@ -127,6 +141,12 @@ func TestMDMAppleCommander(t *testing.T) {
|
||||
|
||||
host := &fleet.Host{ID: 1, UUID: "A", Platform: "darwin"}
|
||||
cmdUUID = uuid.New().String()
|
||||
|
||||
// Mock GetPendingLockCommand to return nil (no pending command)
|
||||
mdmStorage.GetPendingLockCommandFunc = func(ctx context.Context, hostUUID string) (*mdm.Command, string, error) {
|
||||
return nil, "", nil
|
||||
}
|
||||
|
||||
mdmStorage.EnqueueDeviceLockCommandFunc = func(ctx context.Context, gotHost *fleet.Host, cmd *mdm.Command, pin string) error {
|
||||
require.NotNil(t, gotHost)
|
||||
require.Equal(t, host.ID, gotHost.ID)
|
||||
@@ -161,6 +181,269 @@ func TestMDMAppleCommander(t *testing.T) {
|
||||
mdmStorage.RetrievePushInfoFuncInvoked = false
|
||||
}
|
||||
|
||||
func TestMDMAppleCommanderConcurrentDeviceLock(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mdmStorage := &mdmmock.MDMAppleStore{}
|
||||
pushFactory, _ := newMockAPNSPushProviderFactory()
|
||||
pusher := nanomdm_pushsvc.New(
|
||||
mdmStorage,
|
||||
mdmStorage,
|
||||
pushFactory,
|
||||
stdlogfmt.New(),
|
||||
)
|
||||
cmdr := NewMDMAppleCommander(mdmStorage, pusher)
|
||||
|
||||
host := &fleet.Host{ID: 1, UUID: "TEST-HOST", Platform: "darwin"}
|
||||
|
||||
// Variables to track calls (with mutex for thread safety)
|
||||
var mu sync.Mutex
|
||||
var pendingCommand *mdm.Command
|
||||
var pendingPIN string
|
||||
enqueueCalls := 0
|
||||
getPendingCalls := 0
|
||||
|
||||
// Mock GetPendingLockCommand
|
||||
// Need to track state across concurrent calls
|
||||
var commandCreated bool
|
||||
mdmStorage.GetPendingLockCommandFunc = func(ctx context.Context, hostUUID string) (*mdm.Command, string, error) {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
getPendingCalls++
|
||||
require.Equal(t, host.UUID, hostUUID)
|
||||
// After the first command is enqueued, return it as pending
|
||||
if commandCreated && pendingCommand != nil {
|
||||
return pendingCommand, pendingPIN, nil
|
||||
}
|
||||
return nil, "", nil
|
||||
}
|
||||
|
||||
// Mock EnqueueDeviceLockCommand
|
||||
mdmStorage.EnqueueDeviceLockCommandFunc = func(ctx context.Context, gotHost *fleet.Host, cmd *mdm.Command, pin string) error {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
enqueueCalls++
|
||||
require.NotNil(t, gotHost)
|
||||
require.Equal(t, host.ID, gotHost.ID)
|
||||
require.Equal(t, host.UUID, gotHost.UUID)
|
||||
require.Equal(t, "DeviceLock", cmd.Command.RequestType)
|
||||
require.Len(t, pin, 6)
|
||||
// Store the first command as pending, reject others with conflict
|
||||
if !commandCreated {
|
||||
pendingCommand = cmd
|
||||
pendingPIN = pin
|
||||
commandCreated = true
|
||||
return nil
|
||||
}
|
||||
// Command already exists, return conflict error
|
||||
return &mockConflictError{msg: "host already has a pending lock command"}
|
||||
}
|
||||
|
||||
// Mock RetrievePushInfo
|
||||
mdmStorage.RetrievePushInfoFunc = func(ctx context.Context, tokens []string) (map[string]*mdm.Push, error) {
|
||||
res := make(map[string]*mdm.Push)
|
||||
for _, token := range tokens {
|
||||
res[token] = &mdm.Push{
|
||||
PushMagic: "magic",
|
||||
Token: []byte("token"),
|
||||
Topic: "topic",
|
||||
}
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
// Mock RetrievePushCert
|
||||
mdmStorage.RetrievePushCertFunc = func(ctx context.Context, topic string) (*tls.Certificate, string, error) {
|
||||
// Return a mock certificate
|
||||
return &tls.Certificate{}, "staleToken", nil
|
||||
}
|
||||
|
||||
// Mock IsPushCertStale - return false (cert is not stale)
|
||||
mdmStorage.IsPushCertStaleFunc = func(ctx context.Context, topic string, staleToken string) (bool, error) {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// Simulate concurrent lock requests
|
||||
numGoroutines := 10
|
||||
results := make(chan string, numGoroutines)
|
||||
errors := make(chan error, numGoroutines)
|
||||
|
||||
for i := 0; i < numGoroutines; i++ {
|
||||
go func(idx int) {
|
||||
cmdUUID := fmt.Sprintf("cmd-uuid-%d", idx)
|
||||
pin, err := cmdr.DeviceLock(ctx, host, cmdUUID)
|
||||
if err != nil {
|
||||
errors <- err
|
||||
} else {
|
||||
results <- pin
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
|
||||
// Collect results
|
||||
var pins []string
|
||||
for i := 0; i < numGoroutines; i++ {
|
||||
select {
|
||||
case pin := <-results:
|
||||
pins = append(pins, pin)
|
||||
case err := <-errors:
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Verify results
|
||||
require.Len(t, pins, numGoroutines, "All requests should succeed")
|
||||
|
||||
// All PINs should be the same
|
||||
firstPIN := pins[0]
|
||||
for _, pin := range pins {
|
||||
require.Equal(t, firstPIN, pin, "All requests should return the same PIN")
|
||||
}
|
||||
|
||||
// Due to race conditions, multiple goroutines may attempt to enqueue
|
||||
// but only one should succeed, the rest should get conflict errors.
|
||||
// The important thing is that all requests return the same PIN
|
||||
require.GreaterOrEqual(t, enqueueCalls, 1, "At least one enqueue attempt should be made")
|
||||
require.LessOrEqual(t, enqueueCalls, numGoroutines, "At most numGoroutines enqueue attempts")
|
||||
|
||||
// GetPendingLockCommand should have been called multiple times
|
||||
// This includes both initial checks and post-conflict checks
|
||||
require.GreaterOrEqual(t, getPendingCalls, numGoroutines, "GetPendingLockCommand should be called at least once per request")
|
||||
}
|
||||
|
||||
func TestMDMAppleCommanderDeviceLockPushNotificationFailure(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mdmStorage := &mdmmock.MDMAppleStore{}
|
||||
|
||||
// Create a mock push provider that will fail
|
||||
pushProvider := &svcmock.APNSPushProvider{}
|
||||
pushFactory := &svcmock.APNSPushProviderFactory{}
|
||||
pushFactory.NewPushProviderFunc = func(*tls.Certificate) (push.PushProvider, error) {
|
||||
return pushProvider, nil
|
||||
}
|
||||
|
||||
pusher := nanomdm_pushsvc.New(
|
||||
mdmStorage,
|
||||
mdmStorage,
|
||||
pushFactory,
|
||||
stdlogfmt.New(),
|
||||
)
|
||||
cmdr := NewMDMAppleCommander(mdmStorage, pusher)
|
||||
|
||||
host := &fleet.Host{ID: 1, UUID: "TEST-HOST-PUSH-FAIL", Platform: "darwin"}
|
||||
|
||||
// Track whether we're on the first or second request
|
||||
var requestCount int
|
||||
var existingCommand *mdm.Command
|
||||
var existingPIN string
|
||||
|
||||
// Mock GetPendingLockCommand
|
||||
mdmStorage.GetPendingLockCommandFunc = func(ctx context.Context, hostUUID string) (*mdm.Command, string, error) {
|
||||
requestCount++
|
||||
require.Equal(t, host.UUID, hostUUID)
|
||||
|
||||
switch requestCount {
|
||||
case 1:
|
||||
// First request - no pending command
|
||||
return nil, "", nil
|
||||
case 2:
|
||||
// Second request initial check - still no pending command
|
||||
// (hasn't been created yet)
|
||||
return nil, "", nil
|
||||
case 3:
|
||||
// Second request after conflict - return the existing command
|
||||
return existingCommand, existingPIN, nil
|
||||
default:
|
||||
t.Fatalf("Unexpected call to GetPendingLockCommand: %d", requestCount)
|
||||
return nil, "", nil
|
||||
}
|
||||
}
|
||||
|
||||
// Mock EnqueueDeviceLockCommand
|
||||
var enqueueCalls int
|
||||
mdmStorage.EnqueueDeviceLockCommandFunc = func(ctx context.Context, gotHost *fleet.Host, cmd *mdm.Command, pin string) error {
|
||||
enqueueCalls++
|
||||
require.NotNil(t, gotHost)
|
||||
require.Equal(t, host.ID, gotHost.ID)
|
||||
require.Equal(t, "DeviceLock", cmd.Command.RequestType)
|
||||
|
||||
switch enqueueCalls {
|
||||
case 1:
|
||||
// First request succeeds
|
||||
existingCommand = cmd
|
||||
existingPIN = pin
|
||||
return nil
|
||||
case 2:
|
||||
// Second request gets conflict
|
||||
return &mockConflictError{msg: "host already has a pending lock command"}
|
||||
default:
|
||||
t.Fatalf("Unexpected call to EnqueueDeviceLockCommand: %d", enqueueCalls)
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// Mock RetrievePushInfo
|
||||
mdmStorage.RetrievePushInfoFunc = func(ctx context.Context, tokens []string) (map[string]*mdm.Push, error) {
|
||||
res := make(map[string]*mdm.Push)
|
||||
for _, token := range tokens {
|
||||
res[token] = &mdm.Push{
|
||||
PushMagic: "magic",
|
||||
Token: []byte("token"),
|
||||
Topic: "topic",
|
||||
}
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
// Mock RetrievePushCert
|
||||
mdmStorage.RetrievePushCertFunc = func(ctx context.Context, topic string) (*tls.Certificate, string, error) {
|
||||
return &tls.Certificate{}, "staleToken", nil
|
||||
}
|
||||
|
||||
// Mock IsPushCertStale
|
||||
mdmStorage.IsPushCertStaleFunc = func(ctx context.Context, topic string, staleToken string) (bool, error) {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// Configure push provider to fail on conflict scenario
|
||||
var pushAttempts int
|
||||
pushProvider.PushFunc = func(ctx context.Context, pushes []*mdm.Push) (map[string]*push.Response, error) {
|
||||
pushAttempts++
|
||||
|
||||
switch pushAttempts {
|
||||
case 1:
|
||||
// First request - push succeeds
|
||||
return mockSuccessfulPush(ctx, pushes)
|
||||
case 2:
|
||||
// Second request during conflict handling - push fails
|
||||
// This simulates a network error or push service issue
|
||||
return nil, errors.New("push notification service unavailable")
|
||||
default:
|
||||
t.Fatalf("Unexpected push attempt: %d", pushAttempts)
|
||||
return nil, nil
|
||||
}
|
||||
}
|
||||
|
||||
// First request - should succeed normally
|
||||
pin1, err := cmdr.DeviceLock(ctx, host, "cmd-uuid-1")
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, pin1)
|
||||
require.Len(t, pin1, 6)
|
||||
|
||||
// Reset request count for second request
|
||||
requestCount = 1
|
||||
|
||||
// Second concurrent request - should get conflict but still return PIN despite push failure
|
||||
pin2, err := cmdr.DeviceLock(ctx, host, "cmd-uuid-2")
|
||||
require.NoError(t, err, "Should not return error even when push notification fails")
|
||||
require.NotEmpty(t, pin2)
|
||||
require.Equal(t, pin1, pin2, "Should return the same PIN as first request")
|
||||
|
||||
// Verify the expected number of calls
|
||||
require.Equal(t, 2, enqueueCalls, "Should have attempted to enqueue twice")
|
||||
require.Equal(t, 2, pushAttempts, "Should have attempted push twice")
|
||||
require.Equal(t, 3, requestCount, "Should have called GetPendingLockCommand three times")
|
||||
}
|
||||
|
||||
func newMockAPNSPushProviderFactory() (*svcmock.APNSPushProviderFactory, *svcmock.APNSPushProvider) {
|
||||
provider := &svcmock.APNSPushProvider{}
|
||||
provider.PushFunc = mockSuccessfulPush
|
||||
|
||||
@@ -65,6 +65,8 @@ type GetAllMDMConfigAssetsByNameFunc func(ctx context.Context, assetNames []flee
|
||||
|
||||
type GetABMTokenByOrgNameFunc func(ctx context.Context, orgName string) (*fleet.ABMToken, error)
|
||||
|
||||
type GetPendingLockCommandFunc func(ctx context.Context, hostUUID string) (*mdm.Command, string, error)
|
||||
|
||||
type EnqueueDeviceLockCommandFunc func(ctx context.Context, host *fleet.Host, cmd *mdm.Command, pin string) error
|
||||
|
||||
type EnqueueDeviceWipeCommandFunc func(ctx context.Context, host *fleet.Host, cmd *mdm.Command) error
|
||||
@@ -145,6 +147,9 @@ type MDMAppleStore struct {
|
||||
GetABMTokenByOrgNameFunc GetABMTokenByOrgNameFunc
|
||||
GetABMTokenByOrgNameFuncInvoked bool
|
||||
|
||||
GetPendingLockCommandFunc GetPendingLockCommandFunc
|
||||
GetPendingLockCommandFuncInvoked bool
|
||||
|
||||
EnqueueDeviceLockCommandFunc EnqueueDeviceLockCommandFunc
|
||||
EnqueueDeviceLockCommandFuncInvoked bool
|
||||
|
||||
@@ -329,6 +334,13 @@ func (fs *MDMAppleStore) GetABMTokenByOrgName(ctx context.Context, orgName strin
|
||||
return fs.GetABMTokenByOrgNameFunc(ctx, orgName)
|
||||
}
|
||||
|
||||
func (fs *MDMAppleStore) GetPendingLockCommand(ctx context.Context, hostUUID string) (*mdm.Command, string, error) {
|
||||
fs.mu.Lock()
|
||||
fs.GetPendingLockCommandFuncInvoked = true
|
||||
fs.mu.Unlock()
|
||||
return fs.GetPendingLockCommandFunc(ctx, hostUUID)
|
||||
}
|
||||
|
||||
func (fs *MDMAppleStore) EnqueueDeviceLockCommand(ctx context.Context, host *fleet.Host, cmd *mdm.Command, pin string) error {
|
||||
fs.mu.Lock()
|
||||
fs.EnqueueDeviceLockCommandFuncInvoked = true
|
||||
|
||||
@@ -9675,7 +9675,7 @@ func (s *integrationMDMTestSuite) TestLockUnlockWipeWindowsLinux() {
|
||||
require.NotNil(t, getHostResp.Host.MDM.PendingAction)
|
||||
require.Equal(t, string(fleet.PendingActionLock), *getHostResp.Host.MDM.PendingAction)
|
||||
|
||||
// try locking the host while it is pending lock fails
|
||||
// try locking the host while it is pending lock fails for Windows/Linux
|
||||
res := s.DoRaw("POST", fmt.Sprintf("/api/latest/fleet/hosts/%d/lock", host.ID), nil, http.StatusUnprocessableEntity)
|
||||
errMsg := extractServerErrorText(res.Body)
|
||||
require.Contains(t, errMsg, "Host has pending lock request.")
|
||||
@@ -9884,10 +9884,12 @@ func (s *integrationMDMTestSuite) TestLockUnlockWipeMacOS() {
|
||||
require.NotNil(t, getHostResp.Host.MDM.PendingAction)
|
||||
require.Equal(t, "lock", *getHostResp.Host.MDM.PendingAction)
|
||||
|
||||
// try locking the host while it is pending lock fails
|
||||
res := s.DoRaw("POST", fmt.Sprintf("/api/latest/fleet/hosts/%d/lock", host.ID), nil, http.StatusUnprocessableEntity)
|
||||
errMsg := extractServerErrorText(res.Body)
|
||||
require.Contains(t, errMsg, "Host has pending lock request.")
|
||||
// try locking the host while it is pending lock returns the same PIN
|
||||
var lockResp2 lockHostResponse
|
||||
s.DoJSON("POST", fmt.Sprintf("/api/latest/fleet/hosts/%d/lock", host.ID), nil, http.StatusOK, &lockResp2, "view_pin", "true")
|
||||
require.Equal(t, lockResp.UnlockPIN, lockResp2.UnlockPIN, "Should return the same PIN for duplicate lock request")
|
||||
require.Equal(t, fleet.PendingActionLock, lockResp2.PendingAction)
|
||||
require.Equal(t, fleet.DeviceStatusUnlocked, lockResp2.DeviceStatus)
|
||||
|
||||
// simulate a successful MDM result for the lock command
|
||||
cmd, err := mdmClient.Idle()
|
||||
@@ -9907,8 +9909,8 @@ func (s *integrationMDMTestSuite) TestLockUnlockWipeMacOS() {
|
||||
// try to lock the host again
|
||||
s.Do("POST", fmt.Sprintf("/api/latest/fleet/hosts/%d/lock", host.ID), nil, http.StatusConflict)
|
||||
// try to wipe a locked host
|
||||
res = s.Do("POST", fmt.Sprintf("/api/latest/fleet/hosts/%d/wipe", host.ID), nil, http.StatusUnprocessableEntity)
|
||||
errMsg = extractServerErrorText(res.Body)
|
||||
res := s.Do("POST", fmt.Sprintf("/api/latest/fleet/hosts/%d/wipe", host.ID), nil, http.StatusUnprocessableEntity)
|
||||
errMsg := extractServerErrorText(res.Body)
|
||||
require.Contains(t, errMsg, "Host cannot be wiped until it is unlocked.")
|
||||
|
||||
// unlock the host
|
||||
|
||||
Reference in New Issue
Block a user