diff --git a/lib/client/kubesession.go b/lib/client/kubesession.go index 1bf4a8ffc20..0b778fd19fe 100644 --- a/lib/client/kubesession.go +++ b/lib/client/kubesession.go @@ -19,6 +19,7 @@ package client import ( "context" "crypto/tls" + "encoding/json" "fmt" "io" "time" @@ -58,12 +59,28 @@ func NewKubeSession(ctx context.Context, tc *TeleportClient, meta types.SessionT } ws, resp, err := dialer.Dial(joinEndpoint, nil) - defer resp.Body.Close() + if resp != nil && resp.Body != nil { + defer resp.Body.Close() + } if err != nil { - body, _ := io.ReadAll(resp.Body) - fmt.Printf("Handshake failed with status %d\nand body: %v\n", resp.StatusCode, string(body)) cancel() - return nil, trace.Wrap(err) + if resp == nil || resp.Body == nil { + return nil, trace.Wrap(err) + } + + body, _ := io.ReadAll(resp.Body) + var respData map[string]interface{} + if err := json.Unmarshal(body, &respData); err != nil { + return nil, trace.Wrap(err) + } + + if message, ok := respData["message"]; ok { + if message, ok := message.(string); ok { + return nil, trace.Errorf("%v", message) + } + } + + return nil, trace.BadParameter("failed to decode remote error: %v", string(body)) } stream, err := streamproto.NewSessionStream(ws, streamproto.ClientHandshake{Mode: mode}) diff --git a/lib/kube/proxy/forwarder.go b/lib/kube/proxy/forwarder.go index 7b2d97175d9..d27e536449a 100644 --- a/lib/kube/proxy/forwarder.go +++ b/lib/kube/proxy/forwarder.go @@ -21,9 +21,11 @@ import ( "crypto/rand" "crypto/tls" "crypto/x509" + "encoding/json" "encoding/pem" "errors" "fmt" + "io" mathrand "math/rand" "net" "net/http" @@ -850,26 +852,35 @@ func (f *Forwarder) join(ctx *authContext, w http.ResponseWriter, req *http.Requ return nil, trace.Wrap(err) } - stream, err := streamproto.NewSessionStream(ws, streamproto.ServerHandshake{MFARequired: session.PresenceEnabled}) - if err != nil { - return nil, trace.Wrap(err) + if err := func() error { + stream, err := streamproto.NewSessionStream(ws, streamproto.ServerHandshake{MFARequired: session.PresenceEnabled}) + if err != nil { + return trace.Wrap(err) + } + + client := &websocketClientStreams{stream} + party := newParty(*ctx, stream.Mode, client) + go func() { + <-stream.Done() + session.mu.Lock() + defer session.mu.Unlock() + session.leave(party.ID) + }() + + err = session.join(party) + if err != nil { + return trace.Wrap(err) + } + + <-party.closeC + return nil + }(); err != nil { + writeErr := ws.WriteControl(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseInternalServerErr, err.Error()), time.Now().Add(time.Second*10)) + if writeErr != nil { + f.log.WithError(writeErr).Warn("Failed to send early-exit websocket close message.") + } } - client := &websocketClientStreams{stream} - party := newParty(*ctx, stream.Mode, client) - go func() { - <-stream.Done() - session.mu.Lock() - defer session.mu.Unlock() - session.leave(party.ID) - }() - - err = session.join(party) - if err != nil { - return nil, trace.Wrap(err) - } - - <-party.closeC return nil, nil } @@ -888,7 +899,17 @@ func (f *Forwarder) remoteJoin(ctx *authContext, w http.ResponseWriter, req *htt wsTarget, respTarget, err := dialer.Dial(url, nil) if err != nil { - return nil, trace.Wrap(err) + msg, err := io.ReadAll(respTarget.Body) + if err != nil { + return nil, trace.Wrap(err) + } + + var obj map[string]interface{} + if err := json.Unmarshal(msg, &obj); err != nil { + return nil, trace.Wrap(err) + } + + return obj, trace.Wrap(err) } defer wsTarget.Close() defer respTarget.Body.Close() diff --git a/tool/tsh/kube.go b/tool/tsh/kube.go index e3d62906956..b5b515116b0 100644 --- a/tool/tsh/kube.go +++ b/tool/tsh/kube.go @@ -131,7 +131,9 @@ func (c *kubeJoinCommand) run(cf *CLIConf) error { } meta, err := c.getSessionMeta(cf.Context, tc) - if err != nil { + if trace.IsNotFound(err) { + return trace.NotFound("Failed to find session %q. The ID may be incorrect.", c.session) + } else if err != nil { return trace.Wrap(err) } @@ -198,6 +200,7 @@ func (c *kubeJoinCommand) run(cf *CLIConf) error { return trace.Wrap(err) } + tlsConfig.InsecureSkipVerify = cf.InsecureSkipVerify session, err := client.NewKubeSession(cf.Context, tc, meta, tc.KubeProxyAddr, kubeStatus.tlsServerName, types.SessionParticipantMode(c.mode), tlsConfig) if err != nil { return trace.Wrap(err)