diff --git a/server/internal/services/patch_runs.go b/server/internal/services/patch_runs.go index e997cf0..20f9679 100644 --- a/server/internal/services/patch_runs.go +++ b/server/internal/services/patch_runs.go @@ -145,6 +145,46 @@ func setServer(ctx context.Context, runID, serverID, from string, set bson.M) (b return res.MatchedCount > 0, nil } +// setServerForCommand is setServer scoped to one command: a late result from +// an older, already-superseded command (say, a fresh dispatch to the same +// server after a failed send) must not land on the new attempt. +func setServerForCommand(ctx context.Context, runID, serverID, commandID, from string, set bson.M) (bool, error) { + fields := bson.M{} + for k, v := range set { + fields["servers.$."+k] = v + } + res, err := db.Col(patchRunsCol).UpdateOne(ctx, + bson.M{"run_id": runID, "servers": bson.M{"$elemMatch": bson.M{"server_id": serverID, "status": from, "command_id": commandID}}}, + bson.M{"$set": fields}) + if err != nil { + return false, err + } + return res.MatchedCount > 0, nil +} + +// claimServerForDispatch atomically moves one server from queued/waiting_offline +// to patching, guarded on the run still being running and not cancelled: a +// tick that loaded the run before CancelPatchRun wrote cancelled_at must not +// dispatch to it. +func claimServerForDispatch(ctx context.Context, runID, serverID, from, commandID string, now time.Time) (bool, error) { + res, err := db.Col(patchRunsCol).UpdateOne(ctx, + bson.M{ + "run_id": runID, + "status": models.PatchRunRunning, + "cancelled_at": bson.M{"$exists": false}, + "servers": bson.M{"$elemMatch": bson.M{"server_id": serverID, "status": from}}, + }, + bson.M{"$set": bson.M{ + "servers.$.status": models.PatchSrvPatching, + "servers.$.command_id": commandID, + "servers.$.started_at": now, + }}) + if err != nil { + return false, err + } + return res.MatchedCount > 0, nil +} + func loadRun(ctx context.Context, filter bson.M) (*models.PatchRun, error) { var run models.PatchRun err := db.Col(patchRunsCol).FindOne(ctx, filter).Decode(&run) @@ -191,18 +231,34 @@ func advanceRun(ctx context.Context, runID string) { } for _, tr := range patchrun.Advance(*run, now, connected) { if tr.Dispatch { + // Each server's claim and any reset-after-failed-send get their + // own fresh, short-lived context rather than sharing the caller's: + // advanceRun can dispatch to many servers in one tick, each ack + // waiting up to Dispatcher's own timeout, and a caller-supplied + // context (StartPolicyRun's patchCtx, or a leader context that + // AdvancePatchRuns loses partway through) would otherwise expire + // mid-loop and silently drop a later server's writes. cmdID := uuid.New().String() - ok, err := setServer(ctx, runID, tr.ServerID, tr.From, - bson.M{"status": models.PatchSrvPatching, "command_id": cmdID, "started_at": now}) - if err != nil || !ok { + claimCtx, claimCancel := patchCtx() + ok, err := claimServerForDispatch(claimCtx, runID, tr.ServerID, tr.From, cmdID, now) + claimCancel() + if err != nil { + log.Printf("patch run %s: server %s: claim: %v", runID, tr.ServerID, err) + continue + } + if !ok { continue } if err := dispatchPatch(tr.ServerID, cmdID, *run); err != nil { // The agent dropped between the connection check and the send. // Back to waiting: the next tick tries again while the window // is open. - _, _ = setServer(ctx, runID, tr.ServerID, models.PatchSrvPatching, - bson.M{"status": models.PatchSrvWaitingOffline, "command_id": ""}) + resetCtx, resetCancel := patchCtx() + if _, rerr := setServerForCommand(resetCtx, runID, tr.ServerID, cmdID, models.PatchSrvPatching, + bson.M{"status": models.PatchSrvWaitingOffline, "command_id": ""}); rerr != nil { + log.Printf("patch run %s: server %s: reset after failed dispatch: %v", runID, tr.ServerID, rerr) + } + resetCancel() } continue } @@ -305,7 +361,7 @@ func RecordPatchResult(instanceID, serverID string, r *pb.PatchResult) { if updated.FinishedAt != nil { set["finished_at"] = *updated.FinishedAt } - if ok, _ := setServer(ctx, run.RunID, serverID, models.PatchSrvPatching, set); ok && updated.Status == models.PatchSrvRebooting { + if ok, _ := setServerForCommand(ctx, run.RunID, serverID, r.CommandId, models.PatchSrvPatching, set); ok && updated.Status == models.PatchSrvRebooting { LogEvent(instanceID, "patch.reboot", "schedule", serverID, "", fmt.Sprintf("%s rebooting for patch policy %s", s.Hostname, run.PolicyName)) }