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:
Carlo
2025-09-17 15:38:54 -04:00
committed by GitHub
parent ff4cbb3336
commit f25418cae7
10 changed files with 686 additions and 11 deletions
+6 -1
View File
@@ -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{},
+12 -3
View File
@@ -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(
+89
View File
@@ -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")
}
+1
View File
@@ -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
}
+5
View File
@@ -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
+38
View File
@@ -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")
}
+283
View File
@@ -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
+12
View File
@@ -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
+9 -7
View File
@@ -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