From a52c205c34011fcf9c584a82df0816bda090868f Mon Sep 17 00:00:00 2001 From: Nikolay Edigaryev Date: Fri, 21 Jul 2023 00:55:42 +0400 Subject: [PATCH] API(port forward endpoint): handle normal WebSocket closure gracefully (#108) --- internal/controller/api_vms_portforward.go | 75 ++++++++++++++-------- 1 file changed, 47 insertions(+), 28 deletions(-) diff --git a/internal/controller/api_vms_portforward.go b/internal/controller/api_vms_portforward.go index 3d48637..85c7c11 100644 --- a/internal/controller/api_vms_portforward.go +++ b/internal/controller/api_vms_portforward.go @@ -9,6 +9,7 @@ import ( "github.com/cirruslabs/orchard/rpc" "github.com/gin-gonic/gin" "github.com/google/uuid" + "github.com/pkg/errors" "net/http" "nhooyr.io/websocket" "strconv" @@ -41,34 +42,9 @@ func (controller *Controller) portForwardVM(ctx *gin.Context) responder.Responde defer waitContextCancel() // Look-up the VM - var vm *v1.VM - - for { - if lookupResponder := controller.storeView(func(txn storepkg.Transaction) responder.Responder { - vm, err = txn.GetVM(name) - if err != nil { - return responder.Error(err) - } - - return nil - }); lookupResponder != nil { - return lookupResponder - } - - if vm.TerminalState() { - return responder.JSON(http.StatusExpectationFailed, NewErrorResponse("VM is in a terminal state '%s'", vm.Status)) - } - if vm.Status == v1.VMStatusRunning { - // VM is running, proceed - break - } - select { - case <-waitContext.Done(): - return responder.JSON(http.StatusRequestTimeout, NewErrorResponse("VM is not running on '%s' worker", vm.Worker)) - case <-time.After(1 * time.Second): - // try again - continue - } + vm, responderImpl := controller.waitForVM(waitContext, name) + if responderImpl != nil { + return responderImpl } // Request and wait for a connection with a worker @@ -113,6 +89,14 @@ func (controller *Controller) portForwardVM(ctx *gin.Context) responder.Responde wsConnAsNetConn := websocket.NetConn(ctx, wsConn, expectedMsgType) if err := proxy.Connections(wsConnAsNetConn, fromWorkerConnection); err != nil { + var websocketCloseError websocket.CloseError + + // Normal closure from the user + if errors.As(err, &websocketCloseError) && + websocketCloseError.Code == websocket.StatusNormalClosure { + return responder.Empty() + } + controller.logger.Warnf("failed to port-forward: %v", err) } @@ -121,3 +105,38 @@ func (controller *Controller) portForwardVM(ctx *gin.Context) responder.Responde return responder.Error(ctx.Err()) } } + +func (controller *Controller) waitForVM(ctx context.Context, name string) (*v1.VM, responder.Responder) { + var vm *v1.VM + var err error + + for { + if lookupResponder := controller.storeView(func(txn storepkg.Transaction) responder.Responder { + vm, err = txn.GetVM(name) + if err != nil { + return responder.Error(err) + } + + return nil + }); lookupResponder != nil { + return nil, lookupResponder + } + + if vm.TerminalState() { + return nil, responder.JSON(http.StatusExpectationFailed, + NewErrorResponse("VM is in a terminal state '%s'", vm.Status)) + } + if vm.Status == v1.VMStatusRunning { + // VM is running, proceed + return vm, nil + } + select { + case <-ctx.Done(): + return nil, responder.JSON(http.StatusRequestTimeout, + NewErrorResponse("VM is not running on '%s' worker", vm.Worker)) + case <-time.After(1 * time.Second): + // try again + continue + } + } +}