Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
228 changes: 217 additions & 11 deletions pkg/transport/lifecycle.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,13 @@ import (
"context"
"errors"
"io"
"net"
"strings"
"sync"
"time"

"github.com/devsy-org/devsy/pkg/log"
"golang.org/x/crypto/ssh"
)

type CloseReason string
Expand Down Expand Up @@ -38,6 +42,7 @@ const (
TransportSideSSH = SideSSH
TransportSideParent = SideParent
)
const DefaultJoinTimeout = 5 * time.Second

const (
TransportCloseUnknown = CloseUnknown
Expand Down Expand Up @@ -153,6 +158,202 @@ type RunManagedOptions struct {
Handler func(context.Context) error
Metadata LogMetadata
TransportSide Side
JoinTimeout time.Duration
}

type managedOutcome struct {
firstSide Side
parentErr error
handlerErr error
transportErr error
handlerCompleted bool
transportCompleted bool
}

func isTeardownOrCancellationError(err error) bool {
if err == nil {
return false
}
return isCancellationErr(err) || isClosedOrEOFErr(err)
}

func isCancellationErr(err error) bool {
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
return true
}
return strings.Contains(err.Error(), "context canceled")
}

func isClosedOrEOFErr(err error) bool {
return isEOFErr(err) || isClosedNetErr(err)
}

func isEOFErr(err error) bool {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) ||
errors.Is(err, io.ErrClosedPipe) {
return true
}
msg := err.Error()
return strings.Contains(msg, ": EOF") || strings.Contains(msg, "closed pipe")
}

func isClosedNetErr(err error) bool {
if errors.Is(err, net.ErrClosed) {
return true
}
return strings.Contains(err.Error(), "use of closed network connection")
}

func isExitMissingErr(err error) bool {
if err == nil {
return false
}
var exitMissing *ssh.ExitMissingError
if errors.As(err, &exitMissing) {
return true
}
return strings.Contains(err.Error(), "remote command exited without exit status or exit signal")
}

func isTransportTeardownError(err error) bool {
if err == nil {
return true
}
return isTeardownOrCancellationError(err) || isExitMissingErr(err)
}

func isMeaningfulHandlerErr(outcome managedOutcome) bool {
return outcome.handlerCompleted && outcome.handlerErr != nil &&
!isTeardownOrCancellationError(outcome.handlerErr)
}

func isGenuineProviderFailure(outcome managedOutcome) bool {
if !outcome.transportCompleted || outcome.transportErr == nil {
return false
}
return outcome.firstSide == SideProvider && !isTransportTeardownError(outcome.transportErr)
}

func resolveManagedErrors(outcome managedOutcome) error {
if outcome.firstSide == SideParent {
return outcome.parentErr
}
if isMeaningfulHandlerErr(outcome) {
return outcome.handlerErr
}
if isGenuineProviderFailure(outcome) {
return outcome.transportErr
}
if outcome.handlerCompleted && outcome.handlerErr == nil {
return nil
}
if err := resolveParentCancellation(outcome); err != nil {
return err
}
return resolveByFirstSide(outcome)
}

func resolveByFirstSide(outcome managedOutcome) error {
switch outcome.firstSide {
case SideSSH:
return resolveSSHFirst(outcome)
case SideProvider:
return resolveProviderFirst(outcome)
default:
return resolveFallback(outcome)
}
}

func resolveSSHFirst(outcome managedOutcome) error {
if !outcome.handlerCompleted || isTeardownOrCancellationError(outcome.handlerErr) {
if outcome.transportErr != nil {
return outcome.transportErr
}
return nil
}
return outcome.handlerErr
}

func resolveParentCancellation(outcome managedOutcome) error {
if outcome.parentErr != nil &&
(errors.Is(outcome.parentErr, context.Canceled) || errors.Is(outcome.parentErr, context.DeadlineExceeded)) {
if !outcome.handlerCompleted || isTeardownOrCancellationError(outcome.handlerErr) {
return outcome.parentErr
}
}
return nil
}

func resolveProviderFirst(outcome managedOutcome) error {
if !outcome.handlerCompleted || isTeardownOrCancellationError(outcome.handlerErr) {
return outcome.transportErr
}
return outcome.handlerErr
}

func resolveFallback(outcome managedOutcome) error {
if outcome.handlerCompleted && outcome.handlerErr != nil {
return outcome.handlerErr
}
if outcome.transportCompleted && outcome.transportErr != nil {
return outcome.transportErr
}
return outcome.parentErr
}

func waitForFirst(
parent context.Context,
transportSide Side,
handlerDone <-chan error,
connDone <-chan error,
) (managedOutcome, error) {
var outcome managedOutcome
select {
case err := <-connDone:
outcome.firstSide = transportSide
outcome.transportErr = err
outcome.transportCompleted = true
return outcome, err
case err := <-handlerDone:
outcome.firstSide = SideSSH
outcome.handlerErr = err
outcome.handlerCompleted = true
return outcome, err
case <-parent.Done():
outcome.firstSide = SideParent
return outcome, parent.Err()
}
}

func initiateTeardown(conn ManagedConn, firstSide Side, handlerErr error) {
if firstSide == SideSSH && handlerErr == nil {
if cw, ok := conn.(CloseWriter); ok {
if err := cw.CloseWrite(); err == nil {
return
}
}
}
_ = conn.Close()
}

func joinRemaining(
joinCtx context.Context,
handlerDone <-chan error,
connDone <-chan error,
outcome *managedOutcome,
) {
for !outcome.handlerCompleted || !outcome.transportCompleted {
select {
case err := <-handlerDone:
outcome.handlerErr = err
outcome.handlerCompleted = true
case err := <-connDone:
outcome.transportErr = err
outcome.transportCompleted = true
case <-joinCtx.Done():
return
}
}
}

func RunManaged(opts RunManagedOptions) error {
Expand All @@ -165,23 +366,28 @@ func RunManaged(opts RunManagedOptions) error {
if opts.Handler == nil {
return errors.New("handler is required")
}

joinTimeout := opts.JoinTimeout
if joinTimeout <= 0 {
joinTimeout = DefaultJoinTimeout
}

lifecycle, ctx := NewPersistentLifecycle(opts.Parent)
handlerDone := make(chan error, 1)
go func() { handlerDone <- opts.Handler(ctx) }()
connDone := make(chan error, 1)
go func() { connDone <- opts.Conn.Wait() }()

var info CloseInfo
select {
case err := <-connDone:
info = Classify(opts.TransportSide, err, opts.Parent)
case err := <-handlerDone:
info = Classify(SideSSH, err, opts.Parent)
case <-opts.Parent.Done():
info = Classify(SideParent, opts.Parent.Err(), opts.Parent)
}
lifecycle.Close(info)
outcome, firstErr := waitForFirst(opts.Parent, opts.TransportSide, handlerDone, connDone)
lifecycle.Close(Classify(outcome.firstSide, firstErr, opts.Parent))

initiateTeardown(opts.Conn, outcome.firstSide, outcome.handlerErr)
joinCtx, cancelJoin := context.WithTimeout(context.WithoutCancel(opts.Parent), joinTimeout)
defer cancelJoin()
joinRemaining(joinCtx, handlerDone, connDone, &outcome)

_ = opts.Conn.Close()
LogClose(lifecycle.CloseInfo(), opts.Metadata)
return lifecycle.CloseInfo().Err
outcome.parentErr = opts.Parent.Err()
return resolveManagedErrors(outcome)
}
Loading
Loading