diff --git a/manager/orchestrator/restart/restart.go b/manager/orchestrator/restart/restart.go index a13340112b..428acb4749 100644 --- a/manager/orchestrator/restart/restart.go +++ b/manager/orchestrator/restart/restart.go @@ -19,6 +19,25 @@ import ( const defaultOldTaskTimeout = time.Minute +// buffer added to stop grace period to account for SIGKILL and state reporting +const stopGracePeriodBuffer = 5 * time.Second + +func stopGracePeriod(ctx context.Context, t *api.Task) (time.Duration, bool) { + container := t.Spec.GetContainer() + if container == nil || container.StopGracePeriod == nil { + return 0, false + } + grace, err := gogotypes.DurationFromProto(container.StopGracePeriod) + if err != nil { + log.G(ctx).WithError(err).WithField("task.id", t.ID).Error("invalid stop grace period") + return 0, false + } + if grace <= 0 { + return 0, false + } + return grace, true +} + type restartedInstance struct { timestamp time.Time } @@ -485,7 +504,18 @@ func (r *Supervisor) DelayStart(ctx context.Context, _ store.Tx, oldTask *api.Ta close(doneCh) }() - oldTaskTimer := time.NewTimer(r.TaskTimeout) + // if task specifies a StopGracePeriod larger than TaskTimeout, wait for it + // so stop-first doesn't start the new task before the old one stops. + oldTaskTimeout := r.TaskTimeout + if waitForTask { + if grace, ok := stopGracePeriod(ctx, oldTask); ok { + if budget := grace + stopGracePeriodBuffer; budget > oldTaskTimeout { + oldTaskTimeout = budget + } + } + } + + oldTaskTimer := time.NewTimer(oldTaskTimeout) defer oldTaskTimer.Stop() // Wait for the delay to elapse, if one is specified. diff --git a/manager/orchestrator/update/updater_test.go b/manager/orchestrator/update/updater_test.go index 55a9038e7e..a37912d39d 100644 --- a/manager/orchestrator/update/updater_test.go +++ b/manager/orchestrator/update/updater_test.go @@ -702,3 +702,128 @@ func TestUpdaterOrder(t *testing.T) { } } } + +// TestUpdaterStopGracePeriod tests that stop-first updates respect the old task's +// stop grace period before releasing the replacement (#3274). +func TestUpdaterStopGracePeriod(t *testing.T) { + ctx := context.Background() + s := store.NewMemoryStore(nil) + assert.NotNil(t, s) + defer s.Close() + + // simulate slow shutdown: don't progress old task to shutdown yet + watch, cancel := state.Watch(s.WatchQueue(), api.EventUpdateTask{}) + defer cancel() + go func() { + for e := range watch { + task := e.(api.EventUpdateTask).Task + _ = s.Update(func(tx store.Tx) error { + task = store.GetTask(tx, task.ID) + if task == nil { + return nil + } + if task.DesiredState == api.TaskStateRunning && task.Status.State != api.TaskStateRunning { + task.Status.State = api.TaskStateRunning + return store.UpdateTask(tx, task) + } + return nil + }) + } + }() + + const ( + stopGracePeriod = 2 * time.Second + taskTimeout = 50 * time.Millisecond + observeAfter = 500 * time.Millisecond + ) + + service := &api.Service{ + ID: "id1", + Spec: api.ServiceSpec{ + Annotations: api.Annotations{ + Name: "name1", + }, + Task: api.TaskSpec{ + Runtime: &api.TaskSpec_Container{ + Container: &api.ContainerSpec{ + Image: "v:1", + StopGracePeriod: gogotypes.DurationProto(stopGracePeriod), + }, + }, + }, + Mode: &api.ServiceSpec_Replicated{ + Replicated: &api.ReplicatedService{ + Replicas: 1, + }, + }, + Update: &api.UpdateConfig{ + Order: api.UpdateConfig_STOP_FIRST, + Monitor: gogotypes.DurationProto(50 * time.Millisecond), + }, + }, + } + + err := s.Update(func(tx store.Tx) error { + assert.NoError(t, store.CreateService(tx, service)) + task := orchestrator.NewTask(nil, service, 0, "") + task.Status.State = api.TaskStateRunning + assert.NoError(t, store.CreateTask(tx, task)) + return nil + }) + assert.NoError(t, err) + + originalSlots := getRunnableSlotSlice(t, s, service) + require.Len(t, originalSlots, 1) + require.Len(t, originalSlots[0], 1) + oldTaskID := originalSlots[0][0].ID + + service.Spec.Task.GetContainer().Image = "v:2" + updater := NewUpdater(s, restart.NewSupervisor(s), nil, service) + updater.restarts.TaskTimeout = taskTimeout + + runDone := make(chan struct{}) + go func() { + defer close(runDone) + updater.Run(ctx, originalSlots) + }() + + // replacement should still be in READY since grace period has not elapsed + time.Sleep(observeAfter) + + var replacement *api.Task + s.View(func(tx store.ReadTx) { + tasks, err := store.FindTasks(tx, store.ByServiceID(service.ID)) + require.NoError(t, err) + for _, task := range tasks { + if task.ID != oldTaskID { + replacement = task + } + } + }) + require.NotNil(t, replacement, "updater should have created a replacement task") + assert.Equal(t, "v:2", replacement.Spec.GetContainer().Image) + assert.Equal(t, api.TaskStateReady, replacement.DesiredState) + + // finish stopping old task; updater should proceed immediately + err = s.Update(func(tx store.Tx) error { + task := store.GetTask(tx, oldTaskID) + require.NotNil(t, task) + task.Status.State = api.TaskStateShutdown + return store.UpdateTask(tx, task) + }) + assert.NoError(t, err) + + select { + case <-runDone: + case <-time.After(stopGracePeriod): + t.Fatal("updater did not proceed after the old task stopped") + } + + updatedSlots := getRunnableSlotSlice(t, s, service) + require.Len(t, updatedSlots, 1) + for _, slot := range updatedSlots { + for _, task := range slot { + assert.Equal(t, "v:2", task.Spec.GetContainer().Image) + } + } +}