<!-- Add the related story/sub-task/bug number, like Resolves #123, or remove if NA --> **Related issue:** Resolves #45931 # Checklist for submitter If some of the following don't apply, delete the relevant line. - [x] Changes file added for user-visible changes in `changes/`, `orbit/changes/` or `ee/fleetd-chrome/changes`. See [Changes files](https://github.com/fleetdm/fleet/blob/main/docs/Contributing/guides/committing-changes.md#changes-files) for more information. - [x] Input data is properly validated, `SELECT *` is avoided, SQL injection is prevented (using placeholders for values in statements), JS inline code is prevented especially for url redirects, and untrusted data interpolated into shell scripts/commands is validated against shell metacharacters. - [x] Timeouts are implemented and retries are limited to avoid infinite loops - [x] If paths of existing endpoints are modified without backwards compatibility, checked the frontend/CLI for any necessary changes ## Testing - [x] Added/updated automated tests - [x] Where appropriate, [automated tests simulate multiple hosts and test for host isolation](https://github.com/fleetdm/fleet/blob/main/docs/Contributing/reference/patterns-backend.md#unit-testing) (updates to one hosts's records do not affect another) - [x] QA'd all new/changed functionality manually <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **Bug Fixes** * Improved device lock handling so only active pending lock commands are treated as valid. * Fixed stale lock state cases where an old lock reference no longer blocks a new lock request. * When a prior lock command is no longer deliverable, a new lock command is now issued and tracked correctly. * Updated coverage to verify lock status transitions and replacement behavior in these edge cases. <!-- end of auto-generated comment: release notes by coderabbit.ai -->
462 lines
15 KiB
Go
462 lines
15 KiB
Go
package mysql
|
|
|
|
import (
|
|
"context"
|
|
"crypto/tls"
|
|
"database/sql"
|
|
"errors"
|
|
"fmt"
|
|
"log/slog"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
abmctx "github.com/fleetdm/fleet/v4/server/contexts/apple_bm"
|
|
"github.com/fleetdm/fleet/v4/server/contexts/ctxerr"
|
|
"github.com/fleetdm/fleet/v4/server/fleet"
|
|
"github.com/fleetdm/fleet/v4/server/mdm/assets"
|
|
nanodep_client "github.com/fleetdm/fleet/v4/server/mdm/nanodep/client"
|
|
nanodep_mysql "github.com/fleetdm/fleet/v4/server/mdm/nanodep/storage/mysql"
|
|
"github.com/fleetdm/fleet/v4/server/mdm/nanomdm/mdm"
|
|
nanomdm_mysql "github.com/fleetdm/fleet/v4/server/mdm/nanomdm/storage/mysql"
|
|
common_mysql "github.com/fleetdm/fleet/v4/server/platform/mysql"
|
|
"github.com/jmoiron/sqlx"
|
|
)
|
|
|
|
// 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
|
|
}
|
|
|
|
func (e lockConflictError) IsClientError() 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
|
|
|
|
db *sqlx.DB
|
|
logger *slog.Logger
|
|
ds fleet.Datastore
|
|
}
|
|
|
|
// NewMDMAppleMDMStorage returns a MySQL nanomdm storage that uses the Datastore
|
|
// underlying MySQL writer *sql.DB.
|
|
func (ds *Datastore) NewMDMAppleMDMStorage() (*NanoMDMStorage, error) {
|
|
s, err := nanomdm_mysql.New(
|
|
nanomdm_mysql.WithDB(ds.primary.DB),
|
|
nanomdm_mysql.WithLogger(ds.logger),
|
|
nanomdm_mysql.WithReaderFunc(ds.reader),
|
|
)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &NanoMDMStorage{
|
|
MySQLStorage: s,
|
|
db: ds.primary,
|
|
logger: ds.logger,
|
|
ds: ds,
|
|
}, nil
|
|
}
|
|
|
|
// NewTestMDMAppleMDMStorage returns a test MySQL nanomdm storage that uses the
|
|
// Datastore underlying MySQL writer *sql.DB. It allows configuring the async
|
|
// last seen time's capacity and interval and should only be used in tests.
|
|
func (ds *Datastore) NewTestMDMAppleMDMStorage(asyncCap int, asyncInterval time.Duration) (*NanoMDMStorage, error) {
|
|
s, err := nanomdm_mysql.New(
|
|
nanomdm_mysql.WithDB(ds.primary.DB),
|
|
nanomdm_mysql.WithLogger(ds.logger),
|
|
nanomdm_mysql.WithReaderFunc(ds.reader),
|
|
nanomdm_mysql.WithAsyncLastSeen(asyncCap, asyncInterval),
|
|
)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &NanoMDMStorage{
|
|
MySQLStorage: s,
|
|
db: ds.primary,
|
|
logger: ds.logger,
|
|
ds: ds,
|
|
}, nil
|
|
}
|
|
|
|
type pushCertStalenessCheck struct {
|
|
hash string
|
|
updatedAt time.Time
|
|
}
|
|
|
|
// We store staleness check in-memory since it's a short-lived 5 minute time window.
|
|
// And it also means some containers might rotate it faster than 5 minutes depending on the time.
|
|
var (
|
|
pushCertStaleness *pushCertStalenessCheck
|
|
pushCertStalenessMu sync.RWMutex
|
|
)
|
|
|
|
// RetrievePushCert partially implements nanomdm_storage.PushCertStore.
|
|
//
|
|
// Returns the push certificate and its MD5 checksum as the stale token.
|
|
func (s *NanoMDMStorage) RetrievePushCert(
|
|
ctx context.Context, topic string,
|
|
) (*tls.Certificate, string, error) {
|
|
cert, checksum, err := assets.APNSKeyPair(ctx, s.ds)
|
|
if err != nil {
|
|
return nil, "", ctxerr.Wrap(ctx, err, "loading push certificate")
|
|
}
|
|
pushCertStalenessMu.Lock()
|
|
defer pushCertStalenessMu.Unlock()
|
|
checkInMemoryHash(checksum)
|
|
return cert, checksum, nil
|
|
}
|
|
|
|
// checkInMemoryHash checks the incoming hash agains the in-memory hash.
|
|
// if criteria is met, it updates the in-memory hash with the new hash and updatedAt = now.
|
|
func checkInMemoryHash(hash string) {
|
|
if pushCertStaleness == nil || pushCertStaleness.hash != hash || time.Since(pushCertStaleness.updatedAt) > 5*time.Minute {
|
|
// We will not call this unless we are stale, OR on new topic getting a provider, which means we should be fine to update here.
|
|
// Update on new hash, or if it's been more than 5 minutes since last update, to avoid fetching the cert on each stale check.
|
|
pushCertStaleness = &pushCertStalenessCheck{
|
|
hash: hash,
|
|
updatedAt: time.Now(),
|
|
}
|
|
}
|
|
}
|
|
|
|
// IsPushCertStale partially implements nanomdm_storage.PushCertStore.
|
|
//
|
|
// Checks the provided stale token against the in-memory hash of the current push certificate. If they differ, the cert is stale.
|
|
// If the token is the same, it checks if the certificate was last updated more than 5 minutes ago. If so, it re-fetches the certificate and updates the hash for future checks.
|
|
func (s *NanoMDMStorage) IsPushCertStale(ctx context.Context, topic, staleToken string) (bool, error) {
|
|
pushCertStalenessMu.RLock()
|
|
staleness := pushCertStaleness
|
|
pushCertStalenessMu.RUnlock()
|
|
if staleness == nil {
|
|
return true, nil
|
|
}
|
|
if staleness.hash != staleToken {
|
|
s.logger.InfoContext(ctx, "push certificate is stale", "topic", topic, "staleToken", staleToken, "currentHash", staleness.hash, "updatedAt", staleness.updatedAt)
|
|
return true, nil
|
|
}
|
|
|
|
// If updated at is more than 5 minutes ago, re-fetch and re-calculate the has for staleness
|
|
if time.Since(staleness.updatedAt) > 5*time.Minute {
|
|
_, checksum, err := assets.APNSKeyPair(ctx, s.ds)
|
|
if err != nil {
|
|
return false, fmt.Errorf("loading push certificate for staleness check: %w", err)
|
|
}
|
|
pushCertStalenessMu.Lock()
|
|
defer pushCertStalenessMu.Unlock()
|
|
checkInMemoryHash(checksum)
|
|
if checksum != staleToken {
|
|
s.logger.InfoContext(ctx, "push certificate is stale after re-checking", "topic", topic, "staleToken", staleToken, "newHash", checksum)
|
|
return true, nil
|
|
}
|
|
}
|
|
|
|
return false, nil
|
|
}
|
|
|
|
// StorePushCert partially implements nanomdm_storage.PushCertStore.
|
|
func (s *NanoMDMStorage) StorePushCert(ctx context.Context, pemCert, pemKey []byte) error {
|
|
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 neq.active = 1
|
|
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,
|
|
cmd *mdm.Command,
|
|
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)
|
|
|
|
// A non-null lock_ref only blocks a new lock if it still points to a
|
|
// deliverable command. Re-enrollment, SCEP renewal, and wipe flip the
|
|
// queued command to active=0 (see nanomdm ClearQueue), and an inactive
|
|
// command is never sent to the device, so treat it as an orphan ref and
|
|
// let the new lock overwrite it below.
|
|
if err == nil && existingLockRef != nil && *existingLockRef != "" {
|
|
var active bool
|
|
if err := sqlx.GetContext(ctx, tx, &active,
|
|
`SELECT EXISTS(SELECT 1 FROM nano_enrollment_queue WHERE command_uuid = ? AND id = ? AND active = 1)`,
|
|
*existingLockRef, host.UUID); err != nil {
|
|
return ctxerr.Wrap(ctx, err, "checking if existing lock command is active")
|
|
}
|
|
if active {
|
|
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,
|
|
lock_ref,
|
|
unlock_pin,
|
|
fleet_platform
|
|
)
|
|
VALUES (?, ?, ?, ?)
|
|
ON DUPLICATE KEY UPDATE
|
|
wipe_ref = NULL,
|
|
unlock_ref = NULL,
|
|
unlock_pin = VALUES(unlock_pin),
|
|
lock_ref = VALUES(lock_ref)`
|
|
|
|
if _, err := tx.ExecContext(ctx, stmt, host.ID, cmd.CommandUUID, pin, host.FleetPlatform()); err != nil {
|
|
return ctxerr.Wrap(ctx, err, "modifying host_mdm_actions for DeviceLock")
|
|
}
|
|
|
|
return nil
|
|
}, s.logger)
|
|
}
|
|
|
|
func (s *NanoMDMStorage) EnqueueDeviceUnlockCommand(ctx context.Context, host *fleet.Host, cmd *mdm.Command) error {
|
|
return common_mysql.WithRetryTxx(ctx, s.db, func(tx sqlx.ExtContext) error {
|
|
if err := enqueueCommandDB(ctx, tx, []string{host.UUID}, cmd); err != nil {
|
|
return err
|
|
}
|
|
|
|
stmt := `
|
|
INSERT INTO host_mdm_actions (
|
|
host_id,
|
|
unlock_ref,
|
|
fleet_platform
|
|
)
|
|
VALUES (?, ?, ?)
|
|
ON DUPLICATE KEY UPDATE
|
|
unlock_ref = VALUES(unlock_ref),
|
|
unlock_pin = NULL`
|
|
|
|
if _, err := tx.ExecContext(ctx, stmt, host.ID, cmd.CommandUUID, host.FleetPlatform()); err != nil {
|
|
return ctxerr.Wrap(ctx, err, "modifying host_mdm_actions for DeviceUnlock")
|
|
}
|
|
|
|
return nil
|
|
}, s.logger)
|
|
}
|
|
|
|
// EnqueueDeviceWipeCommand enqueues a EraseDevice command for the given host.
|
|
func (s *NanoMDMStorage) EnqueueDeviceWipeCommand(ctx context.Context, host *fleet.Host, cmd *mdm.Command) error {
|
|
return common_mysql.WithRetryTxx(ctx, s.db, func(tx sqlx.ExtContext) error {
|
|
if err := enqueueCommandDB(ctx, tx, []string{host.UUID}, cmd); err != nil {
|
|
return err
|
|
}
|
|
|
|
stmt := `
|
|
INSERT INTO host_mdm_actions (
|
|
host_id,
|
|
wipe_ref,
|
|
fleet_platform
|
|
)
|
|
VALUES (?, ?, ?)
|
|
ON DUPLICATE KEY UPDATE
|
|
wipe_ref = VALUES(wipe_ref)`
|
|
|
|
if _, err := tx.ExecContext(ctx, stmt, host.ID, cmd.CommandUUID, host.FleetPlatform()); err != nil {
|
|
return ctxerr.Wrap(ctx, err, "modifying host_mdm_actions for DeviceWipe")
|
|
}
|
|
|
|
return nil
|
|
}, s.logger)
|
|
}
|
|
|
|
func (s *NanoMDMStorage) GetAllMDMConfigAssetsByName(ctx context.Context, assetNames []fleet.MDMAssetName,
|
|
queryerContext sqlx.QueryerContext,
|
|
) (map[fleet.MDMAssetName]fleet.MDMConfigAsset, error) {
|
|
return s.ds.GetAllMDMConfigAssetsByName(ctx, assetNames, queryerContext)
|
|
}
|
|
|
|
func (s *NanoMDMStorage) GetABMTokenByOrgName(ctx context.Context, orgName string) (*fleet.ABMToken, error) {
|
|
return s.ds.GetABMTokenByOrgName(ctx, orgName)
|
|
}
|
|
|
|
// ExpandEmbeddedSecrets in NanoMDMStorage overrides the implementation in nanomdm_mysql.MySQLStorage.
|
|
func (s *NanoMDMStorage) ExpandEmbeddedSecrets(ctx context.Context, document string) (string, error) {
|
|
return s.ds.ExpandEmbeddedSecrets(ctx, document)
|
|
}
|
|
|
|
// ExpandHostSecrets expands host-scoped secrets in the document using the enrollment ID.
|
|
func (s *NanoMDMStorage) ExpandHostSecrets(ctx context.Context, document string, enrollmentID string) (string, error) {
|
|
return s.ds.ExpandHostSecrets(ctx, document, enrollmentID)
|
|
}
|
|
|
|
func (s *NanoMDMStorage) SetRecoveryLockFailed(ctx context.Context, hostUUID string, errorMsg string) error {
|
|
return s.ds.SetRecoveryLockFailed(ctx, hostUUID, errorMsg)
|
|
}
|
|
|
|
// ClearQueue in NanoMDMStorage overrides the implementation in
|
|
// nanomdm_mysql.MySQLStorage. It does call
|
|
// nanomdm_mysql.MySQLStorage.ClearQueue, but expands on its behavior.
|
|
func (s *NanoMDMStorage) ClearQueue(r *mdm.Request) error {
|
|
err := common_mysql.WithRetryTxx(r.Context, s.db, func(tx sqlx.ExtContext) error {
|
|
if err := s.ds.ClearMDMUpcomingActivitiesDB(r.Context, tx, r.ID); err != nil {
|
|
return err
|
|
}
|
|
return nil
|
|
}, s.logger)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
return s.MySQLStorage.ClearQueue(r)
|
|
}
|
|
|
|
// NewMDMAppleDEPStorage returns a MySQL nanodep storage that uses the Datastore
|
|
// underlying MySQL writer *sql.DB.
|
|
func (ds *Datastore) NewMDMAppleDEPStorage() (*NanoDEPStorage, error) {
|
|
s, err := nanodep_mysql.New(nanodep_mysql.WithDB(ds.primary.DB))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return &NanoDEPStorage{
|
|
MySQLStorage: s,
|
|
ds: ds,
|
|
}, nil
|
|
}
|
|
|
|
// NanoDEPStorage wraps a *nanodep_mysql.MySQLStorage and overrides functionality to load
|
|
// DEP auth tokens from the tables managed by Fleet.
|
|
type NanoDEPStorage struct {
|
|
*nanodep_mysql.MySQLStorage
|
|
ds fleet.Datastore
|
|
}
|
|
|
|
// RetrieveAuthTokens partially implements nanodep.AuthTokensRetriever. NOTE: this method will first
|
|
// check the context for an ABM token; if it doesn't find one, it will fall back to checking the DB.
|
|
// This is so we can use the existing DEP client machinery without major changes. See
|
|
// https://github.com/fleetdm/fleet/issues/21177 for more details.
|
|
func (s *NanoDEPStorage) RetrieveAuthTokens(ctx context.Context, name string) (*nanodep_client.OAuth1Tokens, error) {
|
|
if ctxTok, ok := abmctx.FromContext(ctx); ok {
|
|
return ctxTok, nil
|
|
}
|
|
|
|
token, err := assets.ABMToken(ctx, s.ds, name)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("retrieving token in nano dep storage: %w", err)
|
|
}
|
|
|
|
return token, nil
|
|
}
|
|
|
|
// StoreAuthTokens partially implements nanodep.AuthTokensStorer.
|
|
func (s *NanoDEPStorage) StoreAuthTokens(ctx context.Context, name string, tokens *nanodep_client.OAuth1Tokens) error {
|
|
return errors.New("please use fleet.Datastore to manage MDM assets")
|
|
}
|
|
|
|
func enqueueCommandDB(ctx context.Context, tx sqlx.ExtContext, ids []string, cmd *mdm.Command) error {
|
|
// NOTE: the code to insert into nano_commands and
|
|
// nano_enrollment_queue was copied verbatim from the nanomdm
|
|
// implementation. Ideally we modify some of the interfaces to not
|
|
// duplicate the code here, but that needs more careful planning
|
|
// (which we lack right now)
|
|
if len(ids) < 1 {
|
|
return errors.New("no id(s) supplied to queue command to")
|
|
}
|
|
_, err := tx.ExecContext(
|
|
ctx,
|
|
`INSERT INTO nano_commands (command_uuid, request_type, command) VALUES (?, ?, ?);`,
|
|
cmd.CommandUUID, cmd.Command.RequestType, cmd.Raw,
|
|
)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
query := `INSERT INTO nano_enrollment_queue (id, command_uuid) VALUES (?, ?)`
|
|
query += strings.Repeat(", (?, ?)", len(ids)-1)
|
|
args := make([]interface{}, len(ids)*2)
|
|
for i, id := range ids {
|
|
args[i*2] = id
|
|
args[i*2+1] = cmd.CommandUUID
|
|
}
|
|
if _, err = tx.ExecContext(ctx, query+";", args...); err != nil {
|
|
return err
|
|
}
|
|
|
|
return nil
|
|
}
|