From 78b1ea25d57f6a3d41ce8dbe2c943cf8f4e2e194 Mon Sep 17 00:00:00 2001 From: Luke Kysow <1034429+lkysow@users.noreply.github.com> Date: Mon, 25 May 2020 15:03:08 -0700 Subject: [PATCH] Refactor draining feature - trigger on SIGTERM/INT rather than HTTP POST - remove atlantis drain command - refactor into generic status controller --- cmd/drain.go | 27 --- drain/drain.go | 106 --------- main.go | 4 - server/drain_controller.go | 51 ----- server/drain_controller_test.go | 231 ------------------- server/events/command_runner.go | 28 ++- server/events/command_runner_test.go | 19 +- server/events/drainer.go | 72 +++--- server/events/drainer_test.go | 326 +++++---------------------- server/events/mocks/mock_drainer.go | 12 +- server/events_controller_e2e_test.go | 4 +- server/server.go | 37 ++- server/status_controller.go | 37 +++ server/status_controller_test.go | 82 +++++++ 14 files changed, 262 insertions(+), 774 deletions(-) delete mode 100644 cmd/drain.go delete mode 100644 drain/drain.go delete mode 100644 server/drain_controller.go delete mode 100644 server/drain_controller_test.go create mode 100644 server/status_controller.go create mode 100644 server/status_controller_test.go diff --git a/cmd/drain.go b/cmd/drain.go deleted file mode 100644 index 8269ddc40..000000000 --- a/cmd/drain.go +++ /dev/null @@ -1,27 +0,0 @@ -package cmd - -import ( - "github.com/runatlantis/atlantis/drain" - "github.com/runatlantis/atlantis/server/logging" - "github.com/spf13/cobra" -) - -// DrainCmd performs a drain of the local Atlantis server for all running operations. -// The server itself is not shutdown but drained from all running operations. -// When the command returns, the "atlantis server" process can be stopped securely. -type DrainCmd struct { - Logger *logging.SimpleLogger -} - -// Drain returns the runnable cobra command. -func (v *DrainCmd) Init() *cobra.Command { - return &cobra.Command{ - Use: "drain", - Short: "Perform a drain of the local Atlantis server, waiting for completion before returning", - RunE: func(cmd *cobra.Command, args []string) error { - err := drain.Start(v.Logger) - return err - }, - SilenceErrors: true, - } -} diff --git a/drain/drain.go b/drain/drain.go deleted file mode 100644 index 08f950fd7..000000000 --- a/drain/drain.go +++ /dev/null @@ -1,106 +0,0 @@ -package drain - -import ( - "encoding/json" - "errors" - "io/ioutil" - "net/http" - "time" - - "github.com/runatlantis/atlantis/server" - "github.com/runatlantis/atlantis/server/logging" -) - -// Start begins the shutdown process. -func Start(logger *logging.SimpleLogger) error { - logger.Info("Drain starting") - httpClient := &http.Client{} - resp, err := startDrain(httpClient, logger) - if err != nil { - return err - } - logger.Info("Drain of server initiated succesfully") - for { - if resp.DrainCompleted { - logger.Info("Drain of server completed successfully. You can now send a TERM signal to the server.") - break - } - logger.Info("Drain of server still ongoing, waiting a little bit ...") - time.Sleep(5 * time.Second) - resp, err = getDrainStatus(httpClient, logger) - if err != nil { - return err - } - } - return nil -} - -func startDrain(httpClient *http.Client, logger *logging.SimpleLogger) (*server.DrainResponse, error) { - - req, err := http.NewRequest("POST", "http://localhost:4141/drain", nil) - if err != nil { - logger.Err("Failed to create POST request to /drain endpoint: %s", err) - return nil, err - } - - resp, err := httpClient.Do(req) - if err != nil { - logger.Err("Failed to make POST request to /drain endpoint: %s", err) - return nil, err - } - defer resp.Body.Close() - body, err := ioutil.ReadAll(resp.Body) - if err != nil { - logger.Err("Failed to read reponse body of POST request to /drain endpoint: %s", err) - return nil, err - } - - if resp.StatusCode != http.StatusCreated { - logger.Err("Unexpected status code while making POST request to /drain endpoint: %d", resp.StatusCode) - logger.Info("Response content: %s", string(body)) - return nil, errors.New("Unexpected status code") - } - - var response server.DrainResponse - err = json.Unmarshal(body, &response) - if err != nil { - logger.Err("Failed to parse reponse body of POST request to /drain endpoint: %s", err) - return nil, err - } - return &response, nil -} - -func getDrainStatus(httpClient *http.Client, logger *logging.SimpleLogger) (*server.DrainResponse, error) { - - req, err := http.NewRequest("GET", "http://localhost:4141/drain", nil) - if err != nil { - logger.Err("Failed to create GET request to /drain endpoint: %s", err) - return nil, err - } - - resp, err := httpClient.Do(req) - if err != nil { - logger.Err("Failed to make GET request to /drain endpoint: %s", err) - return nil, err - } - defer resp.Body.Close() - body, err := ioutil.ReadAll(resp.Body) - if err != nil { - logger.Err("Failed to read reponse body of GET request to /drain endpoint: %s", err) - return nil, err - } - - if resp.StatusCode != http.StatusOK { - logger.Err("Unexpected status code while making GET request to /drain endpoint: %d", resp.StatusCode) - logger.Info("Response content: %s", string(body)) - return nil, errors.New("Unexpected status code") - } - - var response server.DrainResponse - err = json.Unmarshal(body, &response) - if err != nil { - logger.Err("Failed to parse reponse body of GET request to /drain endpoint: %s", err) - return nil, err - } - return &response, nil -} diff --git a/main.go b/main.go index 800653820..d2aaea00b 100644 --- a/main.go +++ b/main.go @@ -35,12 +35,8 @@ func main() { } version := &cmd.VersionCmd{AtlantisVersion: atlantisVersion} testdrive := &cmd.TestdriveCmd{} - drainCmd := &cmd.DrainCmd{ - Logger: logging.NewSimpleLogger("cmd", false, logging.Info), - } cmd.RootCmd.AddCommand(server.Init()) cmd.RootCmd.AddCommand(version.Init()) cmd.RootCmd.AddCommand(testdrive.Init()) - cmd.RootCmd.AddCommand(drainCmd.Init()) cmd.Execute() } diff --git a/server/drain_controller.go b/server/drain_controller.go deleted file mode 100644 index a74bade16..000000000 --- a/server/drain_controller.go +++ /dev/null @@ -1,51 +0,0 @@ -package server - -import ( - "encoding/json" - "fmt" - "net/http" - - "github.com/runatlantis/atlantis/server/events" - "github.com/runatlantis/atlantis/server/logging" -) - -// DrainController handles all requests relating to Atlantis drainage (to shutdown properly). -type DrainController struct { - Logger *logging.SimpleLogger - Drainer events.Drainer -} - -type DrainResponse struct { - DrainStarted bool `json:"started"` - DrainCompleted bool `json:"completed"` - OngoingOperationsCounter int `json:"ongoingOperations"` -} - -// Get is the GET /drain route. It renders the current drainage status. -func (d *DrainController) Get(w http.ResponseWriter, r *http.Request) { - d.respondStatus(http.StatusOK, w) -} - -// Post is the POST /drain route. It asks atlantis to finish all ongoing operations and to refuse to start new ones. -func (d *DrainController) Post(w http.ResponseWriter, r *http.Request) { - d.Drainer.StartDrain() - d.respondStatus(http.StatusCreated, w) -} - -func (d *DrainController) respondStatus(responseCode int, w http.ResponseWriter) { - status := d.Drainer.GetStatus() - data, err := json.MarshalIndent(&DrainResponse{ - DrainStarted: status.DrainStarted, - DrainCompleted: status.DrainCompleted, - OngoingOperationsCounter: status.OngoingOperationsCounter, - }, "", " ") - if err != nil { - w.WriteHeader(http.StatusInternalServerError) - fmt.Fprintf(w, "Error creating status json response: %s", err) - return - } - d.Logger.Log(logging.Info, "Drain status: %s", string(data)) - w.WriteHeader(responseCode) - w.Header().Set("Content-Type", "application/json") - w.Write(data) // nolint: errcheck -} diff --git a/server/drain_controller_test.go b/server/drain_controller_test.go deleted file mode 100644 index 5f7760904..000000000 --- a/server/drain_controller_test.go +++ /dev/null @@ -1,231 +0,0 @@ -package server_test - -import ( - "bytes" - "encoding/json" - "io/ioutil" - "net/http" - "net/http/httptest" - "testing" - - "github.com/runatlantis/atlantis/server" - "github.com/runatlantis/atlantis/server/events" - "github.com/runatlantis/atlantis/server/logging" - myTests "github.com/runatlantis/atlantis/testing" -) - -func TestDrainController_Get(t *testing.T) { - type fields struct { - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int - } - type wants struct { - Status int - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int - } - tests := []struct { - name string - fields fields - wants wants - }{ - { - name: "simple", - fields: fields{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 0, - }, - wants: wants{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 0, - Status: http.StatusOK, - }, - }, - { - name: "on ongoing", - fields: fields{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 1, - }, - wants: wants{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 1, - Status: http.StatusOK, - }, - }, - { - name: "started", - fields: fields{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 0, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 0, - Status: http.StatusOK, - }, - }, - { - name: "started and completed", - fields: fields{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - Status: http.StatusOK, - }, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - logger := logging.NewNoopLogger() - r, _ := http.NewRequest("GET", "/drain", bytes.NewBuffer(nil)) - w := httptest.NewRecorder() - dr := &events.SimpleDrainer{ - Logger: logger, - Status: events.DrainStatus{ - DrainStarted: tt.fields.DrainStarted, - DrainCompleted: tt.fields.DrainCompleted, - OngoingOperationsCounter: tt.fields.OngoingOperationsCounter, - }, - } - d := &server.DrainController{ - Logger: logger, - Drainer: dr, - } - d.Get(w, r) - - var result server.DrainResponse - t.Helper() - body, err := ioutil.ReadAll(w.Result().Body) - myTests.Ok(t, err) - myTests.Assert(t, tt.wants.Status == w.Result().StatusCode, "exp %d got %d, body: %s", tt.wants.Status, w.Result().StatusCode, string(body)) - err = json.Unmarshal(body, &result) - myTests.Ok(t, err) - myTests.Assert(t, tt.wants.DrainStarted == result.DrainStarted, "exp %s got %s in DrainStarted of %s", tt.wants.DrainStarted, result.DrainStarted, string(body)) - myTests.Assert(t, tt.wants.DrainCompleted == result.DrainCompleted, "exp %s got %s in DrainCompleted of %s", tt.wants.DrainCompleted, result.DrainCompleted, string(body)) - myTests.Assert(t, tt.wants.OngoingOperationsCounter == result.OngoingOperationsCounter, "exp %s got %s in OngoingOperationsCounter of %s", tt.wants.OngoingOperationsCounter, result.OngoingOperationsCounter, string(body)) - }) - } -} - -func TestDrainController_Post(t *testing.T) { - type fields struct { - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int - } - type wants struct { - Status int - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int - } - tests := []struct { - name string - fields fields - wants wants - }{ - { - name: "simple", - fields: fields{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 0, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - Status: http.StatusCreated, - }, - }, - { - name: "on ongoing", - fields: fields{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 1, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 1, - Status: http.StatusCreated, - }, - }, - { - name: "already started", - fields: fields{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 1, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 1, - Status: http.StatusCreated, - }, - }, - { - name: "already started and completed", - fields: fields{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - Status: http.StatusCreated, - }, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - logger := logging.NewNoopLogger() - r, _ := http.NewRequest("GET", "/drain", bytes.NewBuffer(nil)) - w := httptest.NewRecorder() - dr := &events.SimpleDrainer{ - Logger: logger, - Status: events.DrainStatus{ - DrainStarted: tt.fields.DrainStarted, - DrainCompleted: tt.fields.DrainCompleted, - OngoingOperationsCounter: tt.fields.OngoingOperationsCounter, - }, - } - d := &server.DrainController{ - Logger: logger, - Drainer: dr, - } - d.Post(w, r) - - var result server.DrainResponse - t.Helper() - body, err := ioutil.ReadAll(w.Result().Body) - myTests.Ok(t, err) - myTests.Assert(t, tt.wants.Status == w.Result().StatusCode, "exp %d got %d, body: %s", tt.wants.Status, w.Result().StatusCode, string(body)) - err = json.Unmarshal(body, &result) - myTests.Ok(t, err) - myTests.Assert(t, tt.wants.DrainStarted == result.DrainStarted, "exp %s got %s in DrainStarted of %s", tt.wants.DrainStarted, result.DrainStarted, string(body)) - myTests.Assert(t, tt.wants.DrainCompleted == result.DrainCompleted, "exp %s got %s in DrainCompleted of %s", tt.wants.DrainCompleted, result.DrainCompleted, string(body)) - myTests.Assert(t, tt.wants.OngoingOperationsCounter == result.OngoingOperationsCounter, "exp %s got %s in OngoingOperationsCounter of %s", tt.wants.OngoingOperationsCounter, result.OngoingOperationsCounter, string(body)) - }) - } -} diff --git a/server/events/command_runner.go b/server/events/command_runner.go index 7a7625166..e43d821f4 100644 --- a/server/events/command_runner.go +++ b/server/events/command_runner.go @@ -27,6 +27,10 @@ import ( gitlab "github.com/xanzy/go-gitlab" ) +const ( + ShutdownComment = "Atlantis server is shutting down, please try again later." +) + //go:generate pegomock generate -m --use-experimental-model-gen --package mocks -o mocks/mock_command_runner.go CommandRunner // CommandRunner is the first step after a command request has been parsed. @@ -98,21 +102,18 @@ type DefaultCommandRunner struct { PendingPlanFinder PendingPlanFinder WorkingDir WorkingDir DB *db.BoltDB - Drainer Drainer + Drainer *Drainer } // RunAutoplanCommand runs plan when a pull request is opened or updated. func (c *DefaultCommandRunner) RunAutoplanCommand(baseRepo models.Repo, headRepo models.Repo, pull models.PullRequest, user models.User) { - if canProceed := c.Drainer.TryAddNewOngoingOperation(); !canProceed { - if commentErr := c.VCSClient.CreateComment(baseRepo, pull.Num, "Atlantis server is shutting down, please try again later."); commentErr != nil { - c.Logger.Log(logging.Error, "unable to comment drainage: %s", commentErr) + if opStarted := c.Drainer.StartOp(); !opStarted { + if commentErr := c.VCSClient.CreateComment(baseRepo, pull.Num, ShutdownComment); commentErr != nil { + c.Logger.Log(logging.Error, "unable to comment that Atlantis is shutting down: %s", commentErr) } return } - - defer func() { - c.Drainer.RemoveOngoingOperation() - }() + defer c.Drainer.OpDone() log := c.buildLogger(baseRepo.FullName, pull.Num) defer c.logPanics(baseRepo, pull.Num, log) @@ -176,16 +177,13 @@ func (c *DefaultCommandRunner) RunAutoplanCommand(baseRepo models.Repo, headRepo // the event is further validated before making an additional (potentially // wasteful) call to get the necessary data. func (c *DefaultCommandRunner) RunCommentCommand(baseRepo models.Repo, maybeHeadRepo *models.Repo, maybePull *models.PullRequest, user models.User, pullNum int, cmd *CommentCommand) { - if canProceed := c.Drainer.TryAddNewOngoingOperation(); !canProceed { - if commentErr := c.VCSClient.CreateComment(baseRepo, pullNum, "Atlantis server is shutting down, please try again later."); commentErr != nil { - c.Logger.Log(logging.Error, "unable to comment drainage: %s", commentErr) + if opStarted := c.Drainer.StartOp(); !opStarted { + if commentErr := c.VCSClient.CreateComment(baseRepo, pullNum, ShutdownComment); commentErr != nil { + c.Logger.Log(logging.Error, "unable to comment that Atlantis is shutting down: %s", commentErr) } return } - - defer func() { - c.Drainer.RemoveOngoingOperation() - }() + defer c.Drainer.OpDone() log := c.buildLogger(baseRepo.FullName, pullNum) defer c.logPanics(baseRepo, pullNum, log) diff --git a/server/events/command_runner_test.go b/server/events/command_runner_test.go index 1d76ca7c9..ab0f23a86 100644 --- a/server/events/command_runner_test.go +++ b/server/events/command_runner_test.go @@ -44,7 +44,7 @@ var ch events.DefaultCommandRunner var pullLogger *logging.SimpleLogger var workingDir events.WorkingDir var pendingPlanFinder *mocks.MockPendingPlanFinder -var drainer *mocks.MockDrainer +var drainer *events.Drainer func setup(t *testing.T) *vcsmocks.MockClient { RegisterMockTestingT(t) @@ -60,8 +60,7 @@ func setup(t *testing.T) *vcsmocks.MockClient { projectCommandRunner = mocks.NewMockProjectCommandRunner() workingDir = mocks.NewMockWorkingDir() pendingPlanFinder = mocks.NewMockPendingPlanFinder() - drainer = mocks.NewMockDrainer() - When(drainer.TryAddNewOngoingOperation()).ThenReturn(true) + drainer = &events.Drainer{} When(logger.GetLevel()).ThenReturn(logging.Info) When(logger.NewLogger("runatlantis/atlantis#1", true, logging.Info)). ThenReturn(pullLogger) @@ -89,7 +88,7 @@ func setup(t *testing.T) *vcsmocks.MockClient { func TestRunCommentCommand_LogPanics(t *testing.T) { t.Log("if there is a panic it is commented back on the pull request") vcsClient := setup(t) - When(githubGetter.GetPullRequest(fixtures.GithubRepo, fixtures.Pull.Num)).ThenPanic("OMG PANIC!!!") + When(githubGetter.GetPullRequest(fixtures.GithubRepo, fixtures.Pull.Num)).ThenPanic("panic test - if you're seeing this in a test failure this isn't the failing test") ch.RunCommentCommand(fixtures.GithubRepo, &fixtures.GithubRepo, nil, fixtures.User, 1, &events.CommentCommand{Name: models.PlanCommand}) _, _, comment := vcsClient.VerifyWasCalledOnce().CreateComment(matchers.AnyModelsRepo(), AnyInt(), AnyString()).GetCapturedArguments() Assert(t, strings.Contains(comment, "Error: goroutine panic"), fmt.Sprintf("comment should be about a goroutine panic but was %q", comment)) @@ -244,7 +243,7 @@ func TestRunAutoplanCommand_DeletePlans(t *testing.T) { func TestRunCommentCommand_DrainOngoing(t *testing.T) { t.Log("if drain is ongoing then a message should be displayed") vcsClient := setup(t) - When(drainer.TryAddNewOngoingOperation()).ThenReturn(false) + drainer.ShutdownBlocking() ch.RunCommentCommand(fixtures.GithubRepo, &fixtures.GithubRepo, nil, fixtures.User, fixtures.Pull.Num, nil) vcsClient.VerifyWasCalledOnce().CreateComment(fixtures.GithubRepo, fixtures.Pull.Num, "Atlantis server is shutting down, please try again later.") } @@ -252,16 +251,16 @@ func TestRunCommentCommand_DrainOngoing(t *testing.T) { func TestRunCommentCommand_DrainNotOngoing(t *testing.T) { t.Log("if drain is not ongoing then remove ongoing operation must be called even if panic occured") setup(t) - When(githubGetter.GetPullRequest(fixtures.GithubRepo, fixtures.Pull.Num)).ThenPanic("OMG PANIC!!!") + When(githubGetter.GetPullRequest(fixtures.GithubRepo, fixtures.Pull.Num)).ThenPanic("panic test - if you're seeing this in a test failure this isn't the failing test") ch.RunCommentCommand(fixtures.GithubRepo, &fixtures.GithubRepo, nil, fixtures.User, fixtures.Pull.Num, nil) githubGetter.VerifyWasCalledOnce().GetPullRequest(fixtures.GithubRepo, fixtures.Pull.Num) - drainer.VerifyWasCalledOnce().RemoveOngoingOperation() + Equals(t, 0, drainer.GetStatus().InProgressOps) } func TestRunAutoplanCommand_DrainOngoing(t *testing.T) { t.Log("if drain is ongoing then a message should be displayed") vcsClient := setup(t) - When(drainer.TryAddNewOngoingOperation()).ThenReturn(false) + drainer.ShutdownBlocking() ch.RunAutoplanCommand(fixtures.GithubRepo, fixtures.GithubRepo, fixtures.Pull, fixtures.User) vcsClient.VerifyWasCalledOnce().CreateComment(fixtures.GithubRepo, fixtures.Pull.Num, "Atlantis server is shutting down, please try again later.") } @@ -269,8 +268,8 @@ func TestRunAutoplanCommand_DrainOngoing(t *testing.T) { func TestRunAutoplanCommand_DrainNotOngoing(t *testing.T) { t.Log("if drain is not ongoing then remove ongoing operation must be called even if panic occured") setup(t) - When(projectCommandBuilder.BuildAutoplanCommands(matchers.AnyPtrToEventsCommandContext())).ThenPanic("OMG PANIC!!!") + When(projectCommandBuilder.BuildAutoplanCommands(matchers.AnyPtrToEventsCommandContext())).ThenPanic("panic test - if you're seeing this in a test failure this isn't the failing test") ch.RunAutoplanCommand(fixtures.GithubRepo, fixtures.GithubRepo, fixtures.Pull, fixtures.User) projectCommandBuilder.VerifyWasCalledOnce().BuildAutoplanCommands(matchers.AnyPtrToEventsCommandContext()) - drainer.VerifyWasCalledOnce().RemoveOngoingOperation() + Equals(t, 0, drainer.GetStatus().InProgressOps) } diff --git a/server/events/drainer.go b/server/events/drainer.go index 556305124..ff1fb291d 100644 --- a/server/events/drainer.go +++ b/server/events/drainer.go @@ -2,64 +2,62 @@ package events import ( "sync" - - "github.com/runatlantis/atlantis/server/logging" ) -type Drainer interface { - TryAddNewOngoingOperation() bool - RemoveOngoingOperation() - StartDrain() - GetStatus() DrainStatus -} - -type SimpleDrainer struct { - Logger logging.SimpleLogging - Status DrainStatus +// Drainer is used to gracefully shut down atlantis by waiting for in-progress +// operations to complete. +type Drainer struct { + status DrainStatus mutex sync.Mutex + wg sync.WaitGroup } type DrainStatus struct { - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int + // ShuttingDown is whether we are in the progress of shutting down. + ShuttingDown bool + // InProgressOps is the number of operations currently in progress. + InProgressOps int } -// Try to add an operation as ongoing. Return true if the operation is allowed to start, false if it should be rejected. -func (d *SimpleDrainer) TryAddNewOngoingOperation() bool { +// StartOp tries to start a new operation. It returns false is Atlantis is +// shutting down. +func (d *Drainer) StartOp() bool { d.mutex.Lock() defer d.mutex.Unlock() - if d.Status.DrainStarted { + + if d.status.ShuttingDown { return false } - d.Status.OngoingOperationsCounter++ + d.status.InProgressOps++ + d.wg.Add(1) return true } -// Consider an operation as completed. -func (d *SimpleDrainer) RemoveOngoingOperation() { +// OpDone marks an operation as complete. +func (d *Drainer) OpDone() { d.mutex.Lock() defer d.mutex.Unlock() - d.Status.OngoingOperationsCounter-- - if d.Status.OngoingOperationsCounter < 0 { - d.Logger.Log(logging.Warn, "Drain OngoingOperationsCounter became below 0, this is a bug") - d.Status.OngoingOperationsCounter = 0 - } - if d.Status.DrainStarted && d.Status.OngoingOperationsCounter == 0 { - d.Status.DrainCompleted = true + + d.status.InProgressOps-- + d.wg.Done() + if d.status.InProgressOps < 0 { + // This would be a bug. + d.status.InProgressOps = 0 } } -// Start to drain the server. -func (d *SimpleDrainer) StartDrain() { +// ShutdownBlocking sets "shutting down" to true and blocks until there are no +// in progress operations. +func (d *Drainer) ShutdownBlocking() { + // Set the shutdown status. d.mutex.Lock() - defer d.mutex.Unlock() - d.Status.DrainStarted = true - if d.Status.OngoingOperationsCounter == 0 { - d.Status.DrainCompleted = true - } + d.status.ShuttingDown = true + d.mutex.Unlock() + + // Block until there are no in-progress ops. + d.wg.Wait() } -func (d *SimpleDrainer) GetStatus() DrainStatus { - return d.Status +func (d *Drainer) GetStatus() DrainStatus { + return d.status } diff --git a/server/events/drainer_test.go b/server/events/drainer_test.go index 8db68d6ac..9ae9349a1 100644 --- a/server/events/drainer_test.go +++ b/server/events/drainer_test.go @@ -1,293 +1,69 @@ package events_test import ( + "context" "testing" + "time" "github.com/runatlantis/atlantis/server/events" - "github.com/runatlantis/atlantis/server/logging" - myTests "github.com/runatlantis/atlantis/testing" + . "github.com/runatlantis/atlantis/testing" ) -func TestDrainer_TryAddNewOngoingOperation(t *testing.T) { - type fields struct { - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int - } - type wants struct { - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int - Result bool - } - tests := []struct { - name string - fields fields - wants wants - }{ - { - name: "simple", - fields: fields{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 0, - }, - wants: wants{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 1, - Result: true, - }, - }, - { - name: "already started", - fields: fields{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 1, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 1, - Result: false, - }, - }, - { - name: "already completed", - fields: fields{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 1, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 1, - Result: false, - }, - }, - } +// Test starting and completing ops. +func TestDrainer(t *testing.T) { + d := events.Drainer{} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - logger := logging.NewNoopLogger() - d := &events.SimpleDrainer{ - Logger: logger, - Status: events.DrainStatus{ - DrainStarted: tt.fields.DrainStarted, - DrainCompleted: tt.fields.DrainCompleted, - OngoingOperationsCounter: tt.fields.OngoingOperationsCounter, - }, - } + // Starts at 0. + Equals(t, 0, d.GetStatus().InProgressOps) - result := d.TryAddNewOngoingOperation() + // Add 1. + d.StartOp() + Equals(t, 1, d.GetStatus().InProgressOps) - t.Helper() - myTests.Assert(t, tt.wants.Result == result, "exp %d got %d", tt.wants.Result, result) - myTests.Assert(t, tt.wants.DrainStarted == d.Status.DrainStarted, "exp %s got %s in DrainStarted", tt.wants.DrainStarted, d.Status.DrainStarted) - myTests.Assert(t, tt.wants.DrainCompleted == d.Status.DrainCompleted, "exp %s got %s in DrainCompleted", tt.wants.DrainCompleted, d.Status.DrainCompleted) - myTests.Assert(t, tt.wants.OngoingOperationsCounter == d.Status.OngoingOperationsCounter, "exp %s got %s in OngoingOperationsCounter", tt.wants.OngoingOperationsCounter, d.Status.OngoingOperationsCounter) - }) - } + // Remove 1. + d.OpDone() + Equals(t, 0, d.GetStatus().InProgressOps) + + // Add 2. + d.StartOp() + d.StartOp() + Equals(t, 2, d.GetStatus().InProgressOps) + + // Remove 1. + d.OpDone() + Equals(t, 1, d.GetStatus().InProgressOps) } -func TestDrainer_RemoveOngoingOperation(t *testing.T) { - type fields struct { - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int - } - type wants struct { - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int - } - tests := []struct { - name string - fields fields - wants wants - }{ - { - name: "simple", - fields: fields{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 1, - }, - wants: wants{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 0, - }, - }, - { - name: "already started", - fields: fields{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 1, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - }, - }, - { - name: "going negative - not started", - fields: fields{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 0, - }, - wants: wants{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 0, - }, - }, - { - name: "going negative - started", - fields: fields{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - }, - }, - } +func TestDrainer_Shutdown(t *testing.T) { + d := events.Drainer{} + d.StartOp() - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - logger := logging.NewNoopLogger() - d := &events.SimpleDrainer{ - Logger: logger, - Status: events.DrainStatus{ - DrainStarted: tt.fields.DrainStarted, - DrainCompleted: tt.fields.DrainCompleted, - OngoingOperationsCounter: tt.fields.OngoingOperationsCounter, - }, - } + shutdown := make(chan bool) + go func() { + d.ShutdownBlocking() + close(shutdown) + }() - d.RemoveOngoingOperation() + // Sleep to ensure that ShutdownBlocking has been called. + time.Sleep(300 * time.Millisecond) + + // Starting another op should fail. + Equals(t, false, d.StartOp()) + + // Status should be shutting down. + Equals(t, events.DrainStatus{ + ShuttingDown: true, + InProgressOps: 1, + }, d.GetStatus()) + + // Stop the final operation and wait for shutdown to exit. + d.OpDone() + timer, cancel := context.WithTimeout(context.Background(), 1*time.Second) + defer cancel() + select { + case <-shutdown: + case <-timer.Done(): + Assert(t, false, "Timer reached without shutdown") - t.Helper() - myTests.Assert(t, tt.wants.DrainStarted == d.Status.DrainStarted, "exp %s got %s in DrainStarted", tt.wants.DrainStarted, d.Status.DrainStarted) - myTests.Assert(t, tt.wants.DrainCompleted == d.Status.DrainCompleted, "exp %s got %s in DrainCompleted", tt.wants.DrainCompleted, d.Status.DrainCompleted) - myTests.Assert(t, tt.wants.OngoingOperationsCounter == d.Status.OngoingOperationsCounter, "exp %s got %s in OngoingOperationsCounter", tt.wants.OngoingOperationsCounter, d.Status.OngoingOperationsCounter) - }) } } - -func TestDrainer_StartDrain(t *testing.T) { - type fields struct { - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int - } - type wants struct { - DrainStarted bool - DrainCompleted bool - OngoingOperationsCounter int - } - tests := []struct { - name string - fields fields - wants wants - }{ - { - name: "simple with no ongoing operation", - fields: fields{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 0, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - }, - }, - { - name: "simple with one ongoing operation", - fields: fields{ - DrainStarted: false, - DrainCompleted: false, - OngoingOperationsCounter: 1, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 1, - }, - }, - { - name: "already started", - fields: fields{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 1, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: false, - OngoingOperationsCounter: 1, - }, - }, - { - name: "already started and completed", - fields: fields{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - }, - wants: wants{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 0, - }, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - logger := logging.NewNoopLogger() - d := &events.SimpleDrainer{ - Logger: logger, - Status: events.DrainStatus{ - DrainStarted: tt.fields.DrainStarted, - DrainCompleted: tt.fields.DrainCompleted, - OngoingOperationsCounter: tt.fields.OngoingOperationsCounter, - }, - } - - d.StartDrain() - - t.Helper() - myTests.Assert(t, tt.wants.DrainStarted == d.Status.DrainStarted, "exp %s got %s in DrainStarted", tt.wants.DrainStarted, d.Status.DrainStarted) - myTests.Assert(t, tt.wants.DrainCompleted == d.Status.DrainCompleted, "exp %s got %s in DrainCompleted", tt.wants.DrainCompleted, d.Status.DrainCompleted) - myTests.Assert(t, tt.wants.OngoingOperationsCounter == d.Status.OngoingOperationsCounter, "exp %s got %s in OngoingOperationsCounter", tt.wants.OngoingOperationsCounter, d.Status.OngoingOperationsCounter) - }) - } -} - -func TestDrainer_GetStatus(t *testing.T) { - d := &events.SimpleDrainer{ - Status: events.DrainStatus{ - DrainStarted: true, - DrainCompleted: true, - OngoingOperationsCounter: 12, - }, - } - status := d.GetStatus() - - myTests.Assert(t, d.Status.DrainStarted == status.DrainStarted, "exp %s got %s in DrainStarted", d.Status.DrainStarted, status.DrainStarted) - myTests.Assert(t, d.Status.DrainCompleted == status.DrainCompleted, "exp %s got %s in DrainCompleted", d.Status.DrainCompleted, status.DrainCompleted) - myTests.Assert(t, d.Status.OngoingOperationsCounter == status.OngoingOperationsCounter, "exp %s got %s in OngoingOperationsCounter", d.Status.OngoingOperationsCounter, status.OngoingOperationsCounter) - -} diff --git a/server/events/mocks/mock_drainer.go b/server/events/mocks/mock_drainer.go index b6a1a81df..f99015d64 100644 --- a/server/events/mocks/mock_drainer.go +++ b/server/events/mocks/mock_drainer.go @@ -45,7 +45,7 @@ func (mock *MockDrainer) RemoveOngoingOperation() { panic("mock must not be nil. Use myMock := NewMockDrainer().") } params := []pegomock.Param{} - pegomock.GetGenericMockFrom(mock).Invoke("RemoveOngoingOperation", params, []reflect.Type{}) + pegomock.GetGenericMockFrom(mock).Invoke("OpDone", params, []reflect.Type{}) } func (mock *MockDrainer) StartDrain() { @@ -53,7 +53,7 @@ func (mock *MockDrainer) StartDrain() { panic("mock must not be nil. Use myMock := NewMockDrainer().") } params := []pegomock.Param{} - pegomock.GetGenericMockFrom(mock).Invoke("StartDrain", params, []reflect.Type{}) + pegomock.GetGenericMockFrom(mock).Invoke("ShutdownBlocking", params, []reflect.Type{}) } func (mock *MockDrainer) TryAddNewOngoingOperation() bool { @@ -61,7 +61,7 @@ func (mock *MockDrainer) TryAddNewOngoingOperation() bool { panic("mock must not be nil. Use myMock := NewMockDrainer().") } params := []pegomock.Param{} - result := pegomock.GetGenericMockFrom(mock).Invoke("TryAddNewOngoingOperation", params, []reflect.Type{reflect.TypeOf((*bool)(nil)).Elem()}) + result := pegomock.GetGenericMockFrom(mock).Invoke("StartOp", params, []reflect.Type{reflect.TypeOf((*bool)(nil)).Elem()}) var ret0 bool if len(result) != 0 { if result[0] != nil { @@ -127,7 +127,7 @@ func (c *MockDrainer_GetStatus_OngoingVerification) GetAllCapturedArguments() { func (verifier *VerifierMockDrainer) RemoveOngoingOperation() *MockDrainer_RemoveOngoingOperation_OngoingVerification { params := []pegomock.Param{} - methodInvocations := pegomock.GetGenericMockFrom(verifier.mock).Verify(verifier.inOrderContext, verifier.invocationCountMatcher, "RemoveOngoingOperation", params, verifier.timeout) + methodInvocations := pegomock.GetGenericMockFrom(verifier.mock).Verify(verifier.inOrderContext, verifier.invocationCountMatcher, "OpDone", params, verifier.timeout) return &MockDrainer_RemoveOngoingOperation_OngoingVerification{mock: verifier.mock, methodInvocations: methodInvocations} } @@ -144,7 +144,7 @@ func (c *MockDrainer_RemoveOngoingOperation_OngoingVerification) GetAllCapturedA func (verifier *VerifierMockDrainer) StartDrain() *MockDrainer_StartDrain_OngoingVerification { params := []pegomock.Param{} - methodInvocations := pegomock.GetGenericMockFrom(verifier.mock).Verify(verifier.inOrderContext, verifier.invocationCountMatcher, "StartDrain", params, verifier.timeout) + methodInvocations := pegomock.GetGenericMockFrom(verifier.mock).Verify(verifier.inOrderContext, verifier.invocationCountMatcher, "ShutdownBlocking", params, verifier.timeout) return &MockDrainer_StartDrain_OngoingVerification{mock: verifier.mock, methodInvocations: methodInvocations} } @@ -161,7 +161,7 @@ func (c *MockDrainer_StartDrain_OngoingVerification) GetAllCapturedArguments() { func (verifier *VerifierMockDrainer) TryAddNewOngoingOperation() *MockDrainer_TryAddNewOngoingOperation_OngoingVerification { params := []pegomock.Param{} - methodInvocations := pegomock.GetGenericMockFrom(verifier.mock).Verify(verifier.inOrderContext, verifier.invocationCountMatcher, "TryAddNewOngoingOperation", params, verifier.timeout) + methodInvocations := pegomock.GetGenericMockFrom(verifier.mock).Verify(verifier.inOrderContext, verifier.invocationCountMatcher, "StartOp", params, verifier.timeout) return &MockDrainer_TryAddNewOngoingOperation_OngoingVerification{mock: verifier.mock, methodInvocations: methodInvocations} } diff --git a/server/events_controller_e2e_test.go b/server/events_controller_e2e_test.go index a36ad2eb1..95ad4e1f9 100644 --- a/server/events_controller_e2e_test.go +++ b/server/events_controller_e2e_test.go @@ -422,9 +422,7 @@ func setupE2E(t *testing.T, repoDir string) (server.EventsController, *vcsmocks. globalCfg, err = parser.ParseGlobalCfg(expCfgPath, globalCfg) Ok(t, err) } - drainer := &events.SimpleDrainer{ - Logger: logger, - } + drainer := &events.Drainer{} commandRunner := &events.DefaultCommandRunner{ ProjectCommandRunner: &events.DefaultProjectCommandRunner{ Locker: projectLocker, diff --git a/server/server.go b/server/server.go index 1e66c7c78..7cca509dc 100644 --- a/server/server.go +++ b/server/server.go @@ -75,11 +75,12 @@ type Server struct { Locker locking.Locker EventsController *EventsController LocksController *LocksController - DrainController *DrainController + StatusController *StatusController IndexTemplate TemplateWriter LockDetailTemplate TemplateWriter SSLCertFile string SSLKeyFile string + Drainer *events.Drainer } // Config holds config for server that isn't passed in by the user. @@ -306,10 +307,8 @@ func NewServer(userConfig UserConfig, config Config) (*Server, error) { DefaultTFVersion: defaultTfVersion, TerraformBinDir: terraformClient.TerraformBinDir(), } - drainer := &events.SimpleDrainer{ - Logger: logger, - } - drainController := &DrainController{ + drainer := &events.Drainer{} + statusController := &StatusController{ Logger: logger, Drainer: drainer, } @@ -416,11 +415,12 @@ func NewServer(userConfig UserConfig, config Config) (*Server, error) { Locker: lockingClient, EventsController: eventsController, LocksController: locksController, - DrainController: drainController, + StatusController: statusController, IndexTemplate: indexTemplate, LockDetailTemplate: lockTemplate, SSLKeyFile: userConfig.SSLKeyFile, SSLCertFile: userConfig.SSLCertFile, + Drainer: drainer, }, nil } @@ -430,8 +430,7 @@ func (s *Server) Start() error { return r.URL.Path == "/" || r.URL.Path == "/index.html" }) s.Router.HandleFunc("/healthz", s.Healthz).Methods("GET") - s.Router.HandleFunc("/drain", s.DrainController.Get).Methods("GET") - s.Router.HandleFunc("/drain", s.DrainController.Post).Methods("POST") + s.Router.HandleFunc("/status", s.StatusController.Get).Methods("GET") s.Router.PathPrefix("/static/").Handler(http.FileServer(&assetfs.AssetFS{Asset: static.Asset, AssetDir: static.AssetDir, AssetInfo: static.AssetInfo})) s.Router.HandleFunc("/events", s.EventsController.Post).Methods("POST") s.Router.HandleFunc("/locks", s.LocksController.DeleteLock).Methods("DELETE").Queries("id", "{id:.*}") @@ -467,7 +466,8 @@ func (s *Server) Start() error { }() <-stop - s.Logger.Warn("Received interrupt. Safely shutting down") + s.Logger.Warn("Received interrupt. Waiting for in-progress operations to complete") + s.waitForDrain() ctx, _ := context.WithTimeout(context.Background(), 5*time.Second) // nolint: vet if err := server.Shutdown(ctx); err != nil { return cli.NewExitError(fmt.Sprintf("while shutting down: %s", err), 1) @@ -475,6 +475,25 @@ func (s *Server) Start() error { return nil } +// waitForDrain blocks until draining is complete. +func (s *Server) waitForDrain() { + drainComplete := make(chan bool, 1) + go func() { + s.Drainer.ShutdownBlocking() + drainComplete <- true + }() + ticker := time.NewTicker(5 * time.Second) + for { + select { + case <-drainComplete: + s.Logger.Info("All in-progress operations complete, shutting down") + return + case <-ticker.C: + s.Logger.Info("Waiting for in-progress operations to complete, current in-progress ops: %d", s.Drainer.GetStatus().InProgressOps) + } + } +} + // Index is the / route. func (s *Server) Index(w http.ResponseWriter, _ *http.Request) { locks, err := s.Locker.List() diff --git a/server/status_controller.go b/server/status_controller.go new file mode 100644 index 000000000..7add51671 --- /dev/null +++ b/server/status_controller.go @@ -0,0 +1,37 @@ +package server + +import ( + "encoding/json" + "fmt" + "net/http" + + "github.com/runatlantis/atlantis/server/events" + "github.com/runatlantis/atlantis/server/logging" +) + +// StatusController handles the status of Atlantis. +type StatusController struct { + Logger *logging.SimpleLogger + Drainer *events.Drainer +} + +type StatusResponse struct { + ShuttingDown bool `json:"shutting_down"` + InProgressOps int `json:"in_progress_operations"` +} + +// Get is the GET /status route. +func (d *StatusController) Get(w http.ResponseWriter, r *http.Request) { + status := d.Drainer.GetStatus() + data, err := json.MarshalIndent(&StatusResponse{ + ShuttingDown: status.ShuttingDown, + InProgressOps: status.InProgressOps, + }, "", " ") + if err != nil { + w.WriteHeader(http.StatusInternalServerError) + fmt.Fprintf(w, "Error creating status json response: %s", err) + return + } + w.Header().Set("Content-Type", "application/json") + w.Write(data) // nolint: errcheck +} diff --git a/server/status_controller_test.go b/server/status_controller_test.go new file mode 100644 index 000000000..9ad3a4fc4 --- /dev/null +++ b/server/status_controller_test.go @@ -0,0 +1,82 @@ +package server_test + +import ( + "bytes" + "encoding/json" + "io/ioutil" + "net/http" + "net/http/httptest" + "testing" + + "github.com/runatlantis/atlantis/server" + "github.com/runatlantis/atlantis/server/events" + "github.com/runatlantis/atlantis/server/logging" + . "github.com/runatlantis/atlantis/testing" +) + +func TestStatusController_Startup(t *testing.T) { + logger := logging.NewNoopLogger() + r, _ := http.NewRequest("GET", "/status", bytes.NewBuffer(nil)) + w := httptest.NewRecorder() + dr := &events.Drainer{} + d := &server.StatusController{ + Logger: logger, + Drainer: dr, + } + d.Get(w, r) + + var result server.StatusResponse + body, err := ioutil.ReadAll(w.Result().Body) + Ok(t, err) + Equals(t, 200, w.Result().StatusCode) + err = json.Unmarshal(body, &result) + Ok(t, err) + Equals(t, false, result.ShuttingDown) + Equals(t, 0, result.InProgressOps) +} + +func TestStatusController_InProgress(t *testing.T) { + logger := logging.NewNoopLogger() + r, _ := http.NewRequest("GET", "/status", bytes.NewBuffer(nil)) + w := httptest.NewRecorder() + dr := &events.Drainer{} + dr.StartOp() + + d := &server.StatusController{ + Logger: logger, + Drainer: dr, + } + d.Get(w, r) + + var result server.StatusResponse + body, err := ioutil.ReadAll(w.Result().Body) + Ok(t, err) + Equals(t, 200, w.Result().StatusCode) + err = json.Unmarshal(body, &result) + Ok(t, err) + Equals(t, false, result.ShuttingDown) + Equals(t, 1, result.InProgressOps) +} + +func TestStatusController_Shutdown(t *testing.T) { + logger := logging.NewNoopLogger() + r, _ := http.NewRequest("GET", "/status", bytes.NewBuffer(nil)) + w := httptest.NewRecorder() + dr := &events.Drainer{} + dr.ShutdownBlocking() + + d := &server.StatusController{ + Logger: logger, + Drainer: dr, + } + d.Get(w, r) + + var result server.StatusResponse + body, err := ioutil.ReadAll(w.Result().Body) + Ok(t, err) + Equals(t, 200, w.Result().StatusCode) + err = json.Unmarshal(body, &result) + Ok(t, err) + Equals(t, true, result.ShuttingDown) + Equals(t, 0, result.InProgressOps) +}