From 904e056a041173275a2b5cb7058f256e1ce0ee4a Mon Sep 17 00:00:00 2001 From: Dante Catalfamo <43040593+dantecatalfamo@users.noreply.github.com> Date: Mon, 11 Aug 2025 15:17:57 -0400 Subject: [PATCH] Cancel batch execution API (#31757) #31532 --- changes/31532-cancel-batch-activity | 1 + server/datastore/mysql/jobs.go | 34 ++++- server/datastore/mysql/scripts.go | 106 ++++++++++++++- server/datastore/mysql/scripts_test.go | 122 ++++++++++++++++++ server/fleet/datastore.go | 6 + server/fleet/service.go | 3 + server/mock/datastore_mock.go | 46 +++++-- server/service/handler.go | 1 + server/service/integration_enterprise_test.go | 62 +++++++++ server/service/scripts.go | 55 ++++++++ 10 files changed, 420 insertions(+), 16 deletions(-) create mode 100644 changes/31532-cancel-batch-activity diff --git a/changes/31532-cancel-batch-activity b/changes/31532-cancel-batch-activity new file mode 100644 index 0000000000..ad1fc0f5cd --- /dev/null +++ b/changes/31532-cancel-batch-activity @@ -0,0 +1 @@ +- Added batch script cancel endpoint diff --git a/server/datastore/mysql/jobs.go b/server/datastore/mysql/jobs.go index 85c045c061..159b8defa9 100644 --- a/server/datastore/mysql/jobs.go +++ b/server/datastore/mysql/jobs.go @@ -84,7 +84,7 @@ LIMIT ? return jobs, nil } -func (ds *Datastore) UpdateJob(ctx context.Context, id uint, job *fleet.Job) (*fleet.Job, error) { +func (ds *Datastore) updateJob(ctx context.Context, tx sqlx.ExtContext, id uint, job *fleet.Job) (*fleet.Job, error) { query := ` UPDATE jobs SET @@ -99,7 +99,7 @@ WHERE if !job.NotBefore.IsZero() { notBefore = &job.NotBefore } - _, err := ds.writer(ctx).ExecContext(ctx, query, job.State, job.Retries, job.Error, notBefore, id) + _, err := tx.ExecContext(ctx, query, job.State, job.Retries, job.Error, notBefore, id) if err != nil { return nil, err } @@ -107,6 +107,10 @@ WHERE return job, nil } +func (ds *Datastore) UpdateJob(ctx context.Context, id uint, job *fleet.Job) (*fleet.Job, error) { + return ds.updateJob(ctx, ds.writer(ctx), id, job) +} + func (ds *Datastore) CleanupWorkerJobs(ctx context.Context, failedSince, completedSince time.Duration) (int64, error) { // using not_before instead of created_at/updated_at to be able to use the // existing index, and the difference between those timestamps will be @@ -132,3 +136,29 @@ func (ds *Datastore) CleanupWorkerJobs(ctx context.Context, failedSince, complet n, _ := res.RowsAffected() return n, nil } + +func (ds *Datastore) GetJob(ctx context.Context, jobID uint) (*fleet.Job, error) { + query := ` + SELECT + id, + created_at, + updated_at, + name, + args, + state, + retries, + error, + not_before + FROM + jobs + WHERE + id = ?` + + job := &fleet.Job{} + + if err := sqlx.GetContext(ctx, ds.reader(ctx), job, query, jobID); err != nil { + return nil, ctxerr.Wrap(ctx, err, "get job") + } + + return job, nil +} diff --git a/server/datastore/mysql/scripts.go b/server/datastore/mysql/scripts.go index 0be2fdadb1..52050043bf 100644 --- a/server/datastore/mysql/scripts.go +++ b/server/datastore/mysql/scripts.go @@ -1962,6 +1962,106 @@ func (ds *Datastore) BatchScheduleScript(ctx context.Context, userID *uint, scri return batchExecID, nil } +func (ds *Datastore) CancelBatchScript(ctx context.Context, executionID string) error { + stmt := ` +SELECT + bahr.host_execution_id, + bahr.host_id +FROM + batch_activity_host_results bahr +LEFT JOIN + host_script_results hsr ON bahr.host_execution_id = hsr.execution_id -- I think? +WHERE + bahr.batch_execution_id = ? +AND + hsr.canceled = 0 +AND + hsr.exit_code IS NULL +AND + bahr.error IS NULL` + + stmtSetCanceled := ` +UPDATE + batch_activities ba +SET + finished_at = NOW(), + status = 'finished', + canceled = 1, + num_canceled = (SELECT COUNT(*) FROM batch_activity_host_results WHERE batch_execution_id = ba.execution_id) +WHERE + ba.execution_id = ?` + + stmtCanceled := ` +UPDATE + batch_activities +SET + canceled = 1 +WHERE + execution_id = ?` + + activity, err := ds.GetBatchActivity(ctx, executionID) + if err != nil { + return ctxerr.Wrap(ctx, err, "getting batch activity") + } + + if activity.Status == fleet.BatchExecutionFinished { + return nil + } + + if err := ds.withRetryTxx(ctx, func(tx sqlx.ExtContext) error { + // If job worker exists, mark it as complete to stop it from running + if jobID := activity.JobID; jobID != nil { + job, err := ds.GetJob(ctx, *jobID) + if err != nil { + return ctxerr.Wrap(ctx, err, "failed to find job associated with batch activity") + } + + job.State = fleet.JobStateSuccess + + if _, err := ds.updateJob(ctx, tx, *jobID, job); err != nil { + return ctxerr.Wrap(ctx, err, "updating batch activity job") + } + } + + if activity.Status == fleet.BatchExecutionStarted { + // If the batch activity has started, we need to cancel anything in progress or queued + toCancel := []struct { + HostExecutionID string `db:"host_execution_id"` + HostID uint `db:"host_id"` + }{} + + if err := sqlx.SelectContext(ctx, tx, &toCancel, stmt, executionID); err != nil { + return ctxerr.Wrap(ctx, err, "selecting hosts to cancel") + } + + for _, host := range toCancel { + if _, err := ds.cancelHostUpcomingActivity(ctx, tx, host.HostID, host.HostExecutionID); err != nil { + return ctxerr.Wrap(ctx, err, "canceling upcoming activity") + } + } + + if _, err := tx.ExecContext(ctx, stmtCanceled, executionID); err != nil { + return ctxerr.Wrap(ctx, err, "setting canceled column") + } + + if err := ds.markActivitiesAsCompleted(ctx, tx); err != nil { + return ctxerr.Wrap(ctx, err, "marking job as complete and summarizing counts") + } + } else { + // The batch activity is scheduled, but not started + if _, err := tx.ExecContext(ctx, stmtSetCanceled, executionID); err != nil { + return ctxerr.Wrap(ctx, err, "setting canceled host count") + } + } + + return nil + }); err != nil { + return ctxerr.Wrap(ctx, err, "cancel batch script db transaction") + } + + return nil +} + func (ds *Datastore) GetBatchActivity(ctx context.Context, executionID string) (*fleet.BatchActivity, error) { const stmt = ` SELECT @@ -2121,7 +2221,7 @@ FROM ( COUNT(IF(hsr.exit_code > 0, 1, NULL)) AS num_errored, COUNT(IF(hsr.canceled = 1 AND hsr.exit_code IS NULL, 1, NULL)) AS num_canceled, ( - COUNT(*) + COUNT(*) - COUNT(bahr.error) - COUNT(IF(hsr.exit_code = 0, 1, NULL)) - COUNT(IF(hsr.exit_code > 0, 1, NULL)) @@ -2199,9 +2299,9 @@ func (ds *Datastore) CountBatchScriptExecutions(ctx context.Context, filter flee stmtExecutions := ` SELECT COUNT(*) -FROM +FROM batch_activities ba -JOIN +JOIN scripts s ON ba.script_id = s.id WHERE diff --git a/server/datastore/mysql/scripts_test.go b/server/datastore/mysql/scripts_test.go index 89da44e5bd..0c89d2af61 100644 --- a/server/datastore/mysql/scripts_test.go +++ b/server/datastore/mysql/scripts_test.go @@ -46,6 +46,7 @@ func TestScripts(t *testing.T) { {"BatchExecute", testBatchExecute}, {"BatchExecuteWithStatus", testBatchExecuteWithStatus}, {"BatchScriptSchedule", testBatchScriptSchedule}, + {"BatchScriptCancel", testBatchScriptCancel}, {"TestMarkActivitiesAsCompleted", testMarkActivitiesAsCompleted}, {"DeleteScriptActivatesNextActivity", testDeleteScriptActivatesNextActivity}, {"BatchSetScriptActivatesNextActivity", testBatchSetScriptActivatesNextActivity}, @@ -2312,6 +2313,127 @@ func testMarkActivitiesAsCompleted(t *testing.T, ds *Datastore) { require.Equal(t, fleet.BatchExecutionStarted, batchActivity2.Status) } +func testBatchScriptCancel(t *testing.T, ds *Datastore) { + ctx := context.Background() + + user := test.NewUser(t, ds, "user1", "user@example.com", true) + + team1, err := ds.NewTeam(ctx, &fleet.Team{Name: "team1"}) + require.NoError(t, err) + + host1 := test.NewHost(t, ds, "host1", "10.0.0.3", "host1key", "host1uuid", time.Now()) + host2 := test.NewHost(t, ds, "host2", "10.0.0.4", "host2key", "host2uuid", time.Now()) + host3 := test.NewHost(t, ds, "host3", "10.0.0.4", "host3key", "host3uuid", time.Now()) + hostTeam1 := test.NewHost(t, ds, "hostTeam1", "10.0.0.5", "hostTeam1key", "hostTeam1uuid", time.Now(), test.WithTeamID(team1.ID)) + + test.SetOrbitEnrollment(t, host1, ds) + test.SetOrbitEnrollment(t, host2, ds) + test.SetOrbitEnrollment(t, host3, ds) + test.SetOrbitEnrollment(t, hostTeam1, ds) + + script, err := ds.NewScript(ctx, &fleet.Script{ + Name: "script1.sh", + ScriptContents: "echo hi", + }) + require.NoError(t, err) + + //// + // Immediate execution + // + execID1, err := ds.BatchExecuteScript(ctx, &user.ID, script.ID, []uint{host1.ID, host2.ID}) + require.NoError(t, err) + require.NotEmpty(t, execID1) + + summary1, err := ds.ListBatchScriptExecutions(ctx, fleet.BatchExecutionStatusFilter{ExecutionID: &execID1}) + require.NoError(t, err) + require.Len(t, summary1, 1) + require.Equal(t, fleet.BatchExecutionStarted, summary1[0].Status) + require.False(t, summary1[0].Canceled) + require.Equal(t, uint(2), *summary1[0].NumTargeted) + require.Equal(t, uint(2), *summary1[0].NumPending) + + upcoming1, err := ds.listUpcomingHostScriptExecutions(ctx, host1.ID, false, false) + require.NoError(t, err) + require.Len(t, upcoming1, 1) + + _, _, err = ds.SetHostScriptExecutionResult(ctx, &fleet.HostScriptResultPayload{ + HostID: host1.ID, + ExecutionID: upcoming1[0].ExecutionID, + Output: "", + ExitCode: 0, + }) + require.NoError(t, err) + + upcoming1, err = ds.listUpcomingHostScriptExecutions(ctx, host2.ID, false, false) + require.NoError(t, err) + require.Len(t, upcoming1, 1) + + err = ds.CancelBatchScript(ctx, execID1) + require.NoError(t, err) + + summary1, err = ds.ListBatchScriptExecutions(ctx, fleet.BatchExecutionStatusFilter{ExecutionID: &execID1}) + require.NoError(t, err) + require.Len(t, summary1, 1) + require.Equal(t, fleet.BatchExecutionFinished, summary1[0].Status) + require.True(t, summary1[0].Canceled) + require.Equal(t, uint(0), *summary1[0].NumPending) + require.Equal(t, uint(1), *summary1[0].NumRan) + require.Equal(t, uint(2), *summary1[0].NumTargeted) + require.Equal(t, uint(1), *summary1[0].NumCanceled) + + upcoming1, err = ds.listUpcomingHostScriptExecutions(ctx, host1.ID, false, false) + require.NoError(t, err) + require.Len(t, upcoming1, 0) + + upcoming1, err = ds.listUpcomingHostScriptExecutions(ctx, host2.ID, false, false) + require.NoError(t, err) + require.Len(t, upcoming1, 0) + + //// + // Future execution + // + execID2, err := ds.BatchScheduleScript(ctx, &user.ID, script.ID, []uint{host1.ID, host2.ID}, time.Now().Add(2*time.Hour)) + require.NoError(t, err) + require.NotEmpty(t, execID2) + + summary2, err := ds.ListBatchScriptExecutions(ctx, fleet.BatchExecutionStatusFilter{ExecutionID: &execID2}) + require.NoError(t, err) + require.Len(t, summary2, 1) + require.Equal(t, fleet.BatchExecutionScheduled, summary2[0].Status) + require.False(t, summary2[0].Canceled) + require.Equal(t, uint(2), *summary2[0].NumTargeted) + require.Equal(t, uint(2), *summary2[0].NumPending) + + upcoming2, err := ds.listUpcomingHostScriptExecutions(ctx, host1.ID, false, false) + require.NoError(t, err) + require.Len(t, upcoming2, 0) + + upcoming2, err = ds.listUpcomingHostScriptExecutions(ctx, host2.ID, false, false) + require.NoError(t, err) + require.Len(t, upcoming2, 0) + + err = ds.CancelBatchScript(ctx, execID2) + require.NoError(t, err) + + summary2, err = ds.ListBatchScriptExecutions(ctx, fleet.BatchExecutionStatusFilter{ExecutionID: &execID2}) + require.NoError(t, err) + require.Len(t, summary2, 1) + require.Equal(t, fleet.BatchExecutionFinished, summary2[0].Status) + require.True(t, summary2[0].Canceled) + require.Equal(t, uint(0), *summary2[0].NumPending) + require.Equal(t, uint(0), *summary2[0].NumRan) + require.Equal(t, uint(2), *summary2[0].NumCanceled) + require.Equal(t, uint(2), *summary2[0].NumTargeted) + + upcoming2, err = ds.listUpcomingHostScriptExecutions(ctx, host1.ID, false, false) + require.NoError(t, err) + require.Len(t, upcoming2, 0) + + upcoming2, err = ds.listUpcomingHostScriptExecutions(ctx, host2.ID, false, false) + require.NoError(t, err) + require.Len(t, upcoming2, 0) +} + func testDeleteScriptActivatesNextActivity(t *testing.T, ds *Datastore) { ctx := t.Context() u := test.NewUser(t, ds, "Alice", "alice@example.com", true) diff --git a/server/fleet/datastore.go b/server/fleet/datastore.go index bd600304cc..e17e517379 100644 --- a/server/fleet/datastore.go +++ b/server/fleet/datastore.go @@ -1079,6 +1079,9 @@ type Datastore interface { // provided durations. It returns the number of jobs deleted and an error. CleanupWorkerJobs(ctx context.Context, failedSince, completedSince time.Duration) (int64, error) + // GetJob returns a job from the database + GetJob(ctx context.Context, jobID uint) (*Job, error) + /////////////////////////////////////////////////////////////////////////////// // Debug @@ -1827,6 +1830,9 @@ type Datastore interface { // BatchExecuteSummary returns the summary of a batch script execution BatchExecuteSummary(ctx context.Context, executionID string) (*BatchActivity, error) + // CancelBatchScript cancels the execution of a batch script execution + CancelBatchScript(ctx context.Context, executionID string) error + // ListBatchScriptExecutions returns a filtered list of batch script executions, with summaries. ListBatchScriptExecutions(ctx context.Context, filter BatchExecutionStatusFilter) ([]BatchActivity, error) diff --git a/server/fleet/service.go b/server/fleet/service.go index 7926dfc9eb..bd960f102a 100644 --- a/server/fleet/service.go +++ b/server/fleet/service.go @@ -1197,6 +1197,9 @@ type Service interface { BatchScriptExecutionList(ctx context.Context, filter BatchExecutionStatusFilter) ([]BatchActivity, int64, error) + // BatchScriptCancel cancels a batch script execution + BatchScriptCancel(ctx context.Context, batchExecutionID string) error + // Script-based methods (at least for some platforms, MDM-based for others) LockHost(ctx context.Context, hostID uint, viewPIN bool) (unlockPIN string, err error) UnlockHost(ctx context.Context, hostID uint) (unlockPIN string, err error) diff --git a/server/mock/datastore_mock.go b/server/mock/datastore_mock.go index 26230298ac..996601f3b8 100644 --- a/server/mock/datastore_mock.go +++ b/server/mock/datastore_mock.go @@ -771,6 +771,8 @@ type UpdateJobFunc func(ctx context.Context, id uint, job *fleet.Job) (*fleet.Jo type CleanupWorkerJobsFunc func(ctx context.Context, failedSince time.Duration, completedSince time.Duration) (int64, error) +type GetJobFunc func(ctx context.Context, jobID uint) (*fleet.Job, error) + type InnoDBStatusFunc func(ctx context.Context) (string, error) type ProcessListFunc func(ctx context.Context) ([]fleet.MySQLProcess, error) @@ -1171,20 +1173,22 @@ type BatchSetScriptsFunc func(ctx context.Context, tmID *uint, scripts []*fleet. type BatchExecuteScriptFunc func(ctx context.Context, userID *uint, scriptID uint, hostIDs []uint) (string, error) +type BatchScheduleScriptFunc func(ctx context.Context, userID *uint, scriptID uint, hostIDs []uint, notBefore time.Time) (string, error) + +type GetBatchActivityFunc func(ctx context.Context, executionID string) (*fleet.BatchActivity, error) + +type GetBatchActivityHostResultsFunc func(ctx context.Context, executionID string) ([]*fleet.BatchActivityHostResult, error) + type BatchExecuteSummaryFunc func(ctx context.Context, executionID string) (*fleet.BatchActivity, error) +type CancelBatchScriptFunc func(ctx context.Context, executionID string) error + type ListBatchScriptExecutionsFunc func(ctx context.Context, filter fleet.BatchExecutionStatusFilter) ([]fleet.BatchActivity, error) type CountBatchScriptExecutionsFunc func(ctx context.Context, filter fleet.BatchExecutionStatusFilter) (int64, error) type MarkActivitiesAsCompletedFunc func(ctx context.Context) error -type BatchScheduleScriptFunc func(ctx context.Context, userID *uint, scriptID uint, hostIDs []uint, notBefore time.Time) (string, error) - -type GetBatchActivityFunc func(ctx context.Context, executionID string) (*fleet.BatchActivity, error) - -type GetBatchActivityHostResultsFunc func(ctx context.Context, executionID string) ([]*fleet.BatchActivityHostResult, error) - type GetHostLockWipeStatusFunc func(ctx context.Context, host *fleet.Host) (*fleet.HostLockWipeStatus, error) type LockHostViaScriptFunc func(ctx context.Context, request *fleet.HostScriptRequestPayload, hostFleetPlatform string) error @@ -2570,6 +2574,9 @@ type DataStore struct { CleanupWorkerJobsFunc CleanupWorkerJobsFunc CleanupWorkerJobsFuncInvoked bool + GetJobFunc GetJobFunc + GetJobFuncInvoked bool + InnoDBStatusFunc InnoDBStatusFunc InnoDBStatusFuncInvoked bool @@ -3182,13 +3189,16 @@ type DataStore struct { BatchExecuteSummaryFunc BatchExecuteSummaryFunc BatchExecuteSummaryFuncInvoked bool - LastBatchScriptExecutionsFunc ListBatchScriptExecutionsFunc - LastBatchScriptExecutionsFuncInvoked bool + CancelBatchScriptFunc CancelBatchScriptFunc + CancelBatchScriptFuncInvoked bool + + ListBatchScriptExecutionsFunc ListBatchScriptExecutionsFunc + ListBatchScriptExecutionsFuncInvoked bool CountBatchScriptExecutionsFunc CountBatchScriptExecutionsFunc CountBatchScriptExecutionsFuncInvoked bool - MarkActivitiesAsCompletedFunc MarkActivitiesAsCompletedFunc + MarkActivitiesAsCompletedFunc MarkActivitiesAsCompletedFunc MarkActivitiesAsCompletedFuncInvoked bool GetHostLockWipeStatusFunc GetHostLockWipeStatusFunc @@ -6205,6 +6215,13 @@ func (s *DataStore) CleanupWorkerJobs(ctx context.Context, failedSince time.Dura return s.CleanupWorkerJobsFunc(ctx, failedSince, completedSince) } +func (s *DataStore) GetJob(ctx context.Context, jobID uint) (*fleet.Job, error) { + s.mu.Lock() + s.GetJobFuncInvoked = true + s.mu.Unlock() + return s.GetJobFunc(ctx, jobID) +} + func (s *DataStore) InnoDBStatus(ctx context.Context) (string, error) { s.mu.Lock() s.InnoDBStatusFuncInvoked = true @@ -7633,11 +7650,18 @@ func (s *DataStore) BatchExecuteSummary(ctx context.Context, executionID string) return s.BatchExecuteSummaryFunc(ctx, executionID) } +func (s *DataStore) CancelBatchScript(ctx context.Context, executionID string) error { + s.mu.Lock() + s.CancelBatchScriptFuncInvoked = true + s.mu.Unlock() + return s.CancelBatchScriptFunc(ctx, executionID) +} + func (s *DataStore) ListBatchScriptExecutions(ctx context.Context, filter fleet.BatchExecutionStatusFilter) ([]fleet.BatchActivity, error) { s.mu.Lock() - s.LastBatchScriptExecutionsFuncInvoked = true + s.ListBatchScriptExecutionsFuncInvoked = true s.mu.Unlock() - return s.LastBatchScriptExecutionsFunc(ctx, filter) + return s.ListBatchScriptExecutionsFunc(ctx, filter) } func (s *DataStore) CountBatchScriptExecutions(ctx context.Context, filter fleet.BatchExecutionStatusFilter) (int64, error) { diff --git a/server/service/handler.go b/server/service/handler.go index d76854ea30..dafb88a070 100644 --- a/server/service/handler.go +++ b/server/service/handler.go @@ -500,6 +500,7 @@ func attachFleetAPIRoutes(r *mux.Router, svc fleet.Service, config config.FleetC ue.PATCH("/api/_version_/fleet/scripts/{script_id:[0-9]+}", updateScriptEndpoint, updateScriptRequest{}) ue.DELETE("/api/_version_/fleet/scripts/{script_id:[0-9]+}", deleteScriptEndpoint, deleteScriptRequest{}) ue.POST("/api/_version_/fleet/scripts/batch", batchSetScriptsEndpoint, batchSetScriptsRequest{}) + ue.POST("/api/_version_/fleet/scripts/batch/{batch_execution_id:[a-zA-Z0-9-]+}/cancel", batchScriptCancelEndpoint, batchScriptCancelRequest{}) // Deprecated, will remove in favor of batchScriptExecutionStatusEndpoint when batch script details page is ready. ue.GET("/api/_version_/fleet/scripts/batch/summary/{batch_execution_id:[a-zA-Z0-9-]+}", batchScriptExecutionSummaryEndpoint, batchScriptExecutionSummaryRequest{}) ue.GET("/api/_version_/fleet/scripts/batch/{batch_execution_id:[a-zA-Z0-9-]+}", batchScriptExecutionStatusEndpoint, batchScriptExecutionStatusRequest{}) diff --git a/server/service/integration_enterprise_test.go b/server/service/integration_enterprise_test.go index fce466eb05..f64eeb0df9 100644 --- a/server/service/integration_enterprise_test.go +++ b/server/service/integration_enterprise_test.go @@ -6701,6 +6701,68 @@ func (s *integrationEnterpriseTestSuite) TestRunBatchScript() { ) } +func (s *integrationEnterpriseTestSuite) TestCancelBatchScripts() { + t := s.T() + ctx := context.Background() + + host1 := createOrbitEnrolledHost(t, "linux", "host1", s.ds) + host2 := createOrbitEnrolledHost(t, "linux", "host2", s.ds) + host3 := createOrbitEnrolledHost(t, "linux", "host3", s.ds) + host4 := createOrbitEnrolledHost(t, "linux", "host4", s.ds) + + script, err := s.ds.NewScript(ctx, &fleet.Script{ + Name: "script.sh", + ScriptContents: "echo bonjour", + }) + require.NoError(t, err) + + // Immediate execution + var batchRes batchScriptRunResponse + s.DoJSON("POST", "/api/latest/fleet/scripts/run/batch", batchScriptRunRequest{ + ScriptID: script.ID, + HostIDs: []uint{host1.ID, host2.ID}, + }, http.StatusOK, &batchRes) + require.NotEmpty(t, batchRes.BatchExecutionID) + + var batchStatusResp batchScriptExecutionStatusResponse + s.DoJSON("GET", fmt.Sprintf("/api/latest/fleet/scripts/batch/%s", batchRes.BatchExecutionID), nil, http.StatusOK, &batchStatusResp) + require.Equal(t, *batchStatusResp.ScriptID, script.ID) + require.Equal(t, *batchStatusResp.NumTargeted, uint(2)) + require.Equal(t, *batchStatusResp.NumPending, uint(2)) + + var batchCancelResp batchScriptCancelResponse + s.DoJSON("POST", fmt.Sprintf("/api/latest/fleet/scripts/batch/%s/cancel", batchRes.BatchExecutionID), nil, http.StatusOK, &batchCancelResp) + + s.DoJSON("GET", fmt.Sprintf("/api/latest/fleet/scripts/batch/%s", batchRes.BatchExecutionID), nil, http.StatusOK, &batchStatusResp) + require.Equal(t, *batchStatusResp.ScriptID, script.ID) + require.Equal(t, *batchStatusResp.NumTargeted, uint(2)) + require.Equal(t, *batchStatusResp.NumPending, uint(0)) + require.Equal(t, *batchStatusResp.NumCanceled, uint(2)) + + // Future execution + var batchResScheduled batchScriptRunResponse + scheduleTime := time.Now().Add(3 * time.Hour) + s.DoJSON("POST", "/api/latest/fleet/scripts/run/batch", batchScriptRunRequest{ + ScriptID: script.ID, + HostIDs: []uint{host3.ID, host4.ID}, + NotBefore: &scheduleTime, + }, http.StatusOK, &batchResScheduled) + require.NotEmpty(t, batchResScheduled.BatchExecutionID) + + s.DoJSON("GET", fmt.Sprintf("/api/latest/fleet/scripts/batch/%s", batchResScheduled.BatchExecutionID), nil, http.StatusOK, &batchStatusResp) + require.Equal(t, *batchStatusResp.ScriptID, script.ID) + require.Equal(t, *batchStatusResp.NumTargeted, uint(2)) + require.Equal(t, *batchStatusResp.NumPending, uint(2)) + + s.DoJSON("POST", fmt.Sprintf("/api/latest/fleet/scripts/batch/%s/cancel", batchResScheduled.BatchExecutionID), nil, http.StatusOK, &batchCancelResp) + + s.DoJSON("GET", fmt.Sprintf("/api/latest/fleet/scripts/batch/%s", batchResScheduled.BatchExecutionID), nil, http.StatusOK, &batchStatusResp) + require.Equal(t, *batchStatusResp.ScriptID, script.ID) + require.Equal(t, *batchStatusResp.NumTargeted, uint(2)) + require.Equal(t, *batchStatusResp.NumPending, uint(0)) + require.Equal(t, *batchStatusResp.NumCanceled, uint(2)) +} + func (s *integrationEnterpriseTestSuite) TestRunHostSavedScript() { t := s.T() diff --git a/server/service/scripts.go b/server/service/scripts.go index 1f9bb4a3bd..b17e2de23d 100644 --- a/server/service/scripts.go +++ b/server/service/scripts.go @@ -1166,6 +1166,61 @@ func (svc *Service) BatchScriptExecutionSummary(ctx context.Context, batchExecut return summary, nil } +type batchScriptCancelRequest struct { + BatchExecutionID string `url:"batch_execution_id"` +} + +type batchScriptCancelResponse struct { + Err error `json:"error,omitempty"` +} + +func (r batchScriptCancelResponse) Error() error { return r.Err } + +func batchScriptCancelEndpoint(ctx context.Context, request any, svc fleet.Service) (fleet.Errorer, error) { + req := request.(*batchScriptCancelRequest) + if err := svc.BatchScriptCancel(ctx, req.BatchExecutionID); err != nil { + return batchScriptCancelResponse{Err: err}, nil + } + + return batchScriptCancelResponse{}, nil +} + +func (svc *Service) BatchScriptCancel(ctx context.Context, batchExecutionID string) error { + summaryList, err := svc.ds.ListBatchScriptExecutions(ctx, fleet.BatchExecutionStatusFilter{ + ExecutionID: &batchExecutionID, + }) + if err != nil { + return ctxerr.Wrap(ctx, err, "get batch script summary") + } + + // If the list is empty, it means the batch execution does not exist. + if len(summaryList) == 0 { + // If the user can see a no-team script, we can return a 404 because they have global access. + // Otherwise, we return a 403 to avoid leaking info about which IDs exist. + if err := svc.authz.Authorize(ctx, &fleet.Script{}, fleet.ActionRead); err != nil { + return err + } + svc.authz.SkipAuthorization(ctx) + return ctxerr.Wrap(ctx, err, "get batch script status") + } + + if len(summaryList) > 1 { + return ctxerr.Wrap(ctx, fleet.NewInvalidArgumentError("batch_execution_id", "expected a single batch execution status, got multiple")) + } + + summary := (summaryList)[0] + + if err := svc.authz.Authorize(ctx, &fleet.Script{TeamID: summary.TeamID}, fleet.ActionWrite); err != nil { + return err + } + + if err := svc.ds.CancelBatchScript(ctx, batchExecutionID); err != nil { + return ctxerr.Wrap(ctx, err, "canceling batch script") + } + + return nil +} + func (svc *Service) BatchScriptExecutionStatus(ctx context.Context, batchExecutionID string) (*fleet.BatchActivity, error) { summaryList, err := svc.ds.ListBatchScriptExecutions(ctx, fleet.BatchExecutionStatusFilter{ ExecutionID: &batchExecutionID,