Improve K8S session join error propagation (#15242)

This commit is contained in:
Joel
2022-08-12 17:31:11 +00:00
committed by GitHub
parent f2dd75801a
commit 669d32bbed
3 changed files with 65 additions and 24 deletions
+21 -4
View File
@@ -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})
+40 -19
View File
@@ -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()
+4 -1
View File
@@ -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)